from __future__ import annotations

import json
from dataclasses import asdict
from pathlib import Path

import numpy as np
import pandas as pd

from psx_signal.data.schema import DataIssue, REQUIRED_COLUMNS


class DataValidator:
    def __init__(
        self,
        extreme_return_threshold: float = 0.35,
        expected_session_coverage: float = 0.90,
        minimum_history: int = 252,
    ) -> None:
        self.extreme_return_threshold = extreme_return_threshold
        self.expected_session_coverage = expected_session_coverage
        self.minimum_history = minimum_history

    def validate(self, frame: pd.DataFrame, benchmark_dates: pd.Series | None = None) -> list[DataIssue]:
        issues: list[DataIssue] = []
        missing_columns = sorted(set(REQUIRED_COLUMNS) - set(frame.columns))
        if missing_columns:
            return [DataIssue("error", "missing_columns", f"Missing: {', '.join(missing_columns)}")]
        issues.extend(self._duplicate_issues(frame))
        issues.extend(self._row_issues(frame))
        issues.extend(self._return_issues(frame))
        issues.extend(self._history_issues(frame))
        malformed = ~frame["symbol"].astype("string").str.match(r"^[A-Z0-9.\-]+$", na=False)
        if malformed.any():
            issues.append(DataIssue("error", "malformed_symbol", "Malformed symbol", row_count=int(malformed.sum())))
        if "company_name" in frame:
            inconsistent = frame.dropna(subset=["company_name"]).groupby("symbol")["company_name"].nunique()
            inconsistent = inconsistent[inconsistent > 1]
            if len(inconsistent):
                issues.append(DataIssue(
                    "warning", "inconsistent_company_name", "Multiple company names for symbol",
                    row_count=int(len(inconsistent)), details={"symbols": inconsistent.index.tolist()[:20]},
                ))
        for symbol, count in frame.groupby("symbol").size().items():
            if count < self.minimum_history:
                issues.append(DataIssue(
                    "warning", "short_history",
                    f"Only {count} sessions; minimum configured history is {self.minimum_history}",
                    symbol=str(symbol), row_count=int(count),
                ))
        if benchmark_dates is not None:
            issues.extend(self._session_issues(frame, benchmark_dates))
        adjusted_missing = "adjusted_close" not in frame or frame["adjusted_close"].isna().any()
        if adjusted_missing:
            issues.append(DataIssue(
                "warning", "adjustment_unverified",
                "Historical price adjustment is incomplete; backtests across corporate actions may be distorted.",
            ))
        return issues

    @staticmethod
    def _duplicate_issues(frame: pd.DataFrame) -> list[DataIssue]:
        issues: list[DataIssue] = []
        duplicates = frame.duplicated(["symbol", "date"], keep=False)
        if duplicates.any():
            issues.append(DataIssue(
                "error", "duplicate_symbol_date", "Duplicate symbol/date observations",
                row_count=int(duplicates.sum()),
            ))
        return issues

    @staticmethod
    def _row_issues(frame: pd.DataFrame) -> list[DataIssue]:
        checks = {
            "missing_ohlc": frame[["open", "high", "low", "close"]].isna().any(axis=1),
            "missing_volume": frame["volume"].isna(),
            "non_positive_price": (frame[["open", "high", "low", "close"]] <= 0).any(axis=1),
            "negative_volume": frame["volume"] < 0,
            "zero_volume": frame["volume"] == 0,
            "high_below_low": frame["high"] < frame["low"],
            "open_outside_range": (frame["open"] < frame["low"]) | (frame["open"] > frame["high"]),
            "close_outside_range": (frame["close"] < frame["low"]) | (frame["close"] > frame["high"]),
        }
        severity = {"zero_volume": "info"}
        messages = {
            "missing_ohlc": "Missing OHLC values",
            "missing_volume": "Missing volume",
            "non_positive_price": "Zero or negative OHLC price",
            "negative_volume": "Negative volume",
            "zero_volume": "Zero trading volume",
            "high_below_low": "High is below low",
            "open_outside_range": "Open is outside daily high/low",
            "close_outside_range": "Close is outside daily high/low",
        }
        return [
            DataIssue(severity.get(code, "error"), code, messages[code], row_count=int(mask.sum()))
            for code, mask in checks.items() if mask.any()
        ]

    def _return_issues(self, frame: pd.DataFrame) -> list[DataIssue]:
        ordered = frame.sort_values(["symbol", "date"])
        price_column = "adjusted_close" if "adjusted_close" in ordered and ordered["adjusted_close"].notna().all() else "close"
        returns = ordered.groupby("symbol")[price_column].pct_change(fill_method=None)
        mask = returns.abs() > self.extreme_return_threshold
        if not mask.any():
            return []
        return [DataIssue(
            "warning", "extreme_return",
            f"Absolute daily return exceeds {self.extreme_return_threshold:.0%}; inspect corporate actions/data.",
            row_count=int(mask.sum()),
        )]

    @staticmethod
    def _history_issues(frame: pd.DataFrame) -> list[DataIssue]:
        listing_gaps = frame.sort_values(["symbol", "date"]).groupby("symbol")["date"].diff().dt.days
        mask = listing_gaps > 14
        if not mask.any():
            return []
        return [DataIssue("warning", "long_session_gap", "Calendar gap above 14 days", row_count=int(mask.sum()))]

    def _session_issues(self, frame: pd.DataFrame, benchmark_dates: pd.Series) -> list[DataIssue]:
        expected = set(pd.to_datetime(benchmark_dates).dt.normalize())
        issues: list[DataIssue] = []
        if not expected:
            return issues
        for symbol, group in frame.groupby("symbol", sort=False):
            observed = set(group["date"])
            coverage = len(observed & expected) / len(expected)
            if coverage < self.expected_session_coverage:
                issues.append(DataIssue(
                    "warning", "missing_sessions", f"Benchmark-session coverage is {coverage:.1%}",
                    symbol=str(symbol), row_count=len(expected - observed),
                ))
        return issues

    @staticmethod
    def report(issues: list[DataIssue], output: str | Path | None = None) -> dict[str, object]:
        payload = {
            "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),
            "issues": [asdict(issue) for issue in issues],
        }
        if output:
            Path(output).parent.mkdir(parents=True, exist_ok=True)
            Path(output).write_text(json.dumps(payload, indent=2), encoding="utf-8")
        return payload
