from __future__ import annotations

import json
from dataclasses import dataclass
from datetime import UTC, datetime
from pathlib import Path

import numpy as np
import pandas as pd

from psx_signal.config import Settings
from psx_signal.data.artifacts import MarketDataArtifact
from psx_signal.data.manifest import generate_manifest
from psx_signal.data.providers import KarandaazKSE100Provider
from psx_signal.data.storage import DatasetStore, RawArtifactStore


@dataclass(frozen=True)
class BenchmarkImportResult:
    status: str
    rows: int
    start: str | None
    end: str | None
    raw_artifact: str
    validation_errors: int
    validation_warnings: int
    matching_yahoo_sessions: int
    experiment_end: str | None


class BenchmarkImportService:
    def __init__(self, settings: Settings) -> None:
        self.settings = settings
        self.store = DatasetStore(settings.data.root)
        self.raw_store = RawArtifactStore(settings.data.root)

    def run(self, source: str | Path) -> BenchmarkImportResult:
        source_path = Path(source)
        provider = KarandaazKSE100Provider(source_path)
        raw_reference = self.raw_store.preserve(MarketDataArtifact(
            source_path, "indices", "Karandaaz Data Portal: KSE - 100 index; underlying source PSX",
            None, "karandaaz",
        ))
        benchmark = provider.get_index_history()
        benchmark["raw_artifact"] = raw_reference
        current = self.store.read(self.store.kse100_path)
        if not current.empty and "provider" in current and (current["provider"] == "yahoo_finance").any():
            yahoo = current.copy()
            yahoo["benchmark_role"] = "RESEARCH_SECONDARY_VALIDATION"
            DatasetStore._atomic_parquet(yahoo, self.store.kse100_yahoo_path)
        else:
            yahoo = self.store.read(self.store.kse100_yahoo_path)
        quality = self._validate(benchmark, provider)
        if quality["error_count"]:
            self._write_quality(quality)
            return BenchmarkImportResult(
                "FAILED", len(benchmark), self._date(benchmark, "min"), self._date(benchmark, "max"),
                raw_reference, int(quality["error_count"]), int(quality["warning_count"]), 0, None,
            )
        DatasetStore._atomic_parquet(benchmark, self.store.kse100_path)
        comparison, comparison_summary = self._compare(benchmark, yahoo)
        output = Path("reports/task5")
        output.mkdir(parents=True, exist_ok=True)
        comparison.to_csv(output / "kse100_provider_comparison.csv", index=False)
        (output / "kse100_provider_comparison_summary.json").write_text(
            json.dumps(comparison_summary, indent=2), encoding="utf-8"
        )
        self._write_quality(quality)
        equities = self.store.read(self.store.equities_path)
        alignment = self._alignment(equities, benchmark)
        (output / "kse100_session_alignment.json").write_text(
            json.dumps(alignment, indent=2), encoding="utf-8"
        )
        existing_quality_path = self.store.root / "normalized" / "quality-report.json"
        existing_quality = json.loads(existing_quality_path.read_text()) if existing_quality_path.exists() else {}
        manifest_path = self.store.root / "normalized" / "manifest.json"
        old_manifest = json.loads(manifest_path.read_text()) if manifest_path.exists() else {}
        latest_sync = {
            **old_manifest.get("latest_sync", {}),
            "benchmark_imported_at": datetime.now(UTC).isoformat(),
            "benchmark_provider": "karandaaz",
            "benchmark_status": "SUCCESS",
        }
        manifest = generate_manifest(self.store, "yahoo_finance", existing_quality, latest_sync)
        manifest["experiment_end"] = alignment["experiment_end"]
        manifest["benchmark_common_session_coverage"] = alignment["benchmark_coverage_of_equity_sessions"]
        manifest_path.write_text(json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8")
        return BenchmarkImportResult(
            "SUCCESS", len(benchmark), self._date(benchmark, "min"), self._date(benchmark, "max"),
            raw_reference, 0, int(quality["warning_count"]), int(comparison_summary["matching_sessions"]),
            alignment["experiment_end"],
        )

    def _validate(
        self, benchmark: pd.DataFrame, provider: KarandaazKSE100Provider
    ) -> dict[str, object]:
        issues: list[dict[str, object]] = []
        if provider.metadata_rows_skipped:
            issues.append({"severity": "info", "code": "metadata_rows_skipped", "count": provider.metadata_rows_skipped})
        if provider.identical_duplicates_removed:
            issues.append({"severity": "warning", "code": "identical_duplicate_dates", "count": provider.identical_duplicates_removed})
        checks = {
            "null_close": benchmark["close"].isna(),
            "non_positive_close": benchmark["close"] <= 0,
        }
        for code, mask in checks.items():
            if mask.any():
                issues.append({"severity": "error", "code": code, "count": int(mask.sum())})
        if benchmark["date"].duplicated().any():
            issues.append({"severity": "error", "code": "duplicate_dates", "count": int(benchmark["date"].duplicated().sum())})
        if benchmark[["open", "high", "low"]].notna().any().any():
            high_low = benchmark["high"] < benchmark["low"]
            outside = (
                (benchmark["open"] < benchmark["low"]) | (benchmark["open"] > benchmark["high"])
                | (benchmark["close"] < benchmark["low"]) | (benchmark["close"] > benchmark["high"])
            )
            if high_low.any():
                issues.append({"severity": "error", "code": "high_below_low", "count": int(high_low.sum())})
            if outside.any():
                issues.append({"severity": "error", "code": "ohlc_outside_range", "count": int(outside.sum())})
        else:
            issues.append({"severity": "info", "code": "close_only_source", "count": len(benchmark)})
        returns = benchmark.sort_values("date")["close"].pct_change(fill_method=None)
        extreme = returns.abs() > 0.08
        if extreme.any():
            issues.append({
                "severity": "warning", "code": "abnormal_index_return", "count": int(extreme.sum()),
                "maximum_absolute_return": float(returns.abs().max()),
            })
        gaps = benchmark.sort_values("date")["date"].diff().dt.days
        if (gaps > 14).any():
            issues.append({"severity": "warning", "code": "long_gap", "count": int((gaps > 14).sum())})
        weekends = benchmark["date"].dt.dayofweek >= 5
        if weekends.any():
            issues.append({
                "severity": "warning", "code": "weekend_dates", "count": int(weekends.sum()),
                "dates": benchmark.loc[weekends, "date"].dt.strftime("%Y-%m-%d").tolist(),
            })
        equities = self.store.read(self.store.equities_path)
        alignment = self._alignment(equities, benchmark)
        if alignment["benchmark_coverage_of_equity_sessions"] < 0.90:
            issues.append({"severity": "error", "code": "insufficient_session_coverage", **alignment})
        elif alignment["equity_only_dates"] or alignment["benchmark_only_dates"]:
            issues.append({"severity": "warning", "code": "session_misalignment", **alignment})
        stale_days = alignment["equity_days_after_experiment_end"]
        if stale_days:
            issues.append({"severity": "info", "code": "common_period_truncation", "count": stale_days})
        return {
            "valid": not any(issue["severity"] == "error" for issue in issues),
            "error_count": sum(issue["severity"] == "error" for issue in issues),
            "warning_count": sum(issue["severity"] == "warning" for issue in issues),
            "info_count": sum(issue["severity"] == "info" for issue in issues),
            "coverage": {"start": self._date(benchmark, "min"), "end": self._date(benchmark, "max"), "rows": len(benchmark)},
            "provider": "karandaaz", "underlying_source": "PSX",
            "provider_trust": "RESEARCH_SECONDARY", "issues": issues,
        }

    @staticmethod
    def _compare(benchmark: pd.DataFrame, yahoo: pd.DataFrame) -> tuple[pd.DataFrame, dict[str, object]]:
        if yahoo.empty:
            return pd.DataFrame(), {"matching_sessions": 0, "status": "NO_YAHOO_VALIDATION_DATA"}
        left = benchmark[["date", "close"]].rename(columns={"close": "karandaaz_close"})
        right = yahoo[["date", "close"]].rename(columns={"close": "yahoo_close"})
        joined = left.merge(right, on="date", how="inner").sort_values("date")
        joined["absolute_close_difference"] = (joined["karandaaz_close"] - joined["yahoo_close"]).abs()
        joined["percentage_difference"] = joined["absolute_close_difference"] / joined["karandaaz_close"].abs()
        joined["karandaaz_return"] = joined["karandaaz_close"].pct_change(fill_method=None)
        joined["yahoo_return"] = joined["yahoo_close"].pct_change(fill_method=None)
        joined["return_series_difference"] = joined["karandaaz_return"] - joined["yahoo_return"]
        exact = np.isclose(joined["karandaaz_close"], joined["yahoo_close"], rtol=1e-8, atol=1e-8)
        near = joined["percentage_difference"] <= 0.001
        summary = {
            "matching_sessions": len(joined), "exact_values": int(exact.sum()),
            "nearly_exact_within_0_1_percent": int(near.sum()),
            "median_percentage_difference": float(joined["percentage_difference"].median()),
            "p95_percentage_difference": float(joined["percentage_difference"].quantile(0.95)),
            "maximum_percentage_difference": float(joined["percentage_difference"].max()),
            "daily_return_correlation": float(joined["karandaaz_return"].corr(joined["yahoo_return"])),
            "material_over_0_1_percent": int((~near).sum()),
        }
        return joined, summary

    @staticmethod
    def _alignment(equities: pd.DataFrame, benchmark: pd.DataFrame) -> dict[str, object]:
        if equities.empty or benchmark.empty:
            return {
                "experiment_end": None, "equity_sessions": 0, "benchmark_sessions": 0,
                "common_sessions": 0, "equity_only_dates": 0, "benchmark_only_dates": 0,
                "benchmark_coverage_of_equity_sessions": 0.0, "equity_days_after_experiment_end": 0,
            }
        equity_dates = set(pd.to_datetime(equities["date"]).dt.normalize())
        benchmark_dates = set(pd.to_datetime(benchmark["date"]).dt.normalize())
        experiment_start = max(min(equity_dates), min(benchmark_dates))
        experiment_end = min(max(equity_dates), max(benchmark_dates))
        equity_period = {value for value in equity_dates if experiment_start <= value <= experiment_end}
        benchmark_period = {value for value in benchmark_dates if experiment_start <= value <= experiment_end}
        common = equity_period & benchmark_period
        return {
            "experiment_start": str(experiment_start.date()), "experiment_end": str(experiment_end.date()),
            "equity_sessions": len(equity_period), "benchmark_sessions": len(benchmark_period),
            "common_sessions": len(common), "equity_only_dates": len(equity_period - benchmark_period),
            "benchmark_only_dates": len(benchmark_period - equity_period),
            "benchmark_coverage_of_equity_sessions": len(common) / max(len(equity_period), 1),
            "equity_days_after_experiment_end": len([value for value in equity_dates if value > experiment_end]),
            "equity_only_date_sample": [str(value.date()) for value in sorted(equity_period - benchmark_period)[:20]],
            "benchmark_only_date_sample": [str(value.date()) for value in sorted(benchmark_period - equity_period)[:20]],
        }

    @staticmethod
    def _date(frame: pd.DataFrame, operation: str) -> str | None:
        if frame.empty:
            return None
        value = frame["date"].min() if operation == "min" else frame["date"].max()
        return str(pd.Timestamp(value).date())

    @staticmethod
    def _write_quality(quality: dict[str, object]) -> None:
        output = Path("reports/task5/kse100_quality.json")
        output.parent.mkdir(parents=True, exist_ok=True)
        output.write_text(json.dumps(quality, indent=2), encoding="utf-8")
