from __future__ import annotations

import json
from dataclasses import asdict, dataclass, field
from datetime import UTC, date, datetime
from pathlib import Path

import pandas as pd

from psx_signal.config import Settings
from psx_signal.data.artifacts import MarketDataArtifact
from psx_signal.data.diagnostics import detect_discontinuities, reconcile_discontinuities
from psx_signal.data.manifest import generate_manifest
from psx_signal.data.master import build_stock_master, read_stock_metadata
from psx_signal.data.providers.yahoo_provider import YahooDownload, YahooFinanceProvider
from psx_signal.data.schema import CANONICAL_INDEX_COLUMNS, CORPORATE_ACTION_COLUMNS
from psx_signal.data.storage import DataConflictError, DatasetStore, RawArtifactStore
from psx_signal.data.validation import DataValidator


PRICE_COLUMNS = ["open", "high", "low", "close", "volume", "adjusted_close"]


@dataclass
class YahooBootstrapResult:
    status: str
    attempted: int
    successful: int
    failed: int
    rows_added: int
    index_rows_added: int
    action_rows_added: int
    rejected_rows: int = 0
    failures: list[dict[str, str]] = field(default_factory=list)
    manifest: dict[str, object] = field(default_factory=dict)
    download_status_path: str = "reports/yahoo_download_status.csv"
    coverage_path: str = "reports/task4/yahoo_coverage.csv"


class YahooBootstrapService:
    def __init__(self, settings: Settings) -> None:
        self.settings = settings
        self.provider = YahooFinanceProvider(settings.data)
        self.store = DatasetStore(settings.data.root)
        self.raw_store = RawArtifactStore(settings.data.root)
        self.cache_dir = Path(settings.data.root) / "inbox" / "yahoo" / ".cache"

    @staticmethod
    def read_candidates(path: str | Path, limit: int | None = None) -> pd.DataFrame:
        frame = pd.read_csv(path, dtype="string")
        if "symbol" not in frame and "psx_symbol" in frame:
            frame = frame.rename(columns={"psx_symbol": "symbol"})
        if "symbol" not in frame:
            raise ValueError("Candidate file requires symbol or psx_symbol")
        frame["symbol"] = frame["symbol"].str.strip().str.upper()
        frame = frame[frame["symbol"].str.match(r"^[A-Z0-9.\-]+$", na=False)]
        frame = frame.drop_duplicates("symbol", keep="first")
        return frame.head(limit) if limit else frame

    def run(
        self, start: date, end: date, symbols_file: str | Path,
        limit: int | None = None, force: bool = False,
    ) -> YahooBootstrapResult:
        if end < start:
            raise ValueError("end must be on or after start")
        candidates = self.read_candidates(symbols_file, limit)
        self.cache_dir.mkdir(parents=True, exist_ok=True)
        initial_equity_rows = len(self.store.read(self.store.equities_path))
        downloads = [
            self._cached_or_download(symbol, start, end, force)
            for symbol in candidates["symbol"].tolist()
        ]
        status_rows: list[dict[str, object]] = []
        coverage_rows: list[dict[str, object]] = []
        failures: list[dict[str, str]] = []
        successful_symbols: list[str] = []
        rows_added = 0
        actions_added = 0

        for download in downloads:
            try:
                if download.status != "SUCCESS":
                    raise RuntimeError(download.error or "No Yahoo Finance history")
                raw_reference = self._preserve_raw(download, start, end)
                normalized = self.provider.normalize_download(download)
                normalized["raw_artifact"] = raw_reference
                _, added = self.store.upsert(
                    normalized, self.store.equities_path, ["symbol", "date"], PRICE_COLUMNS,
                    ["date", "symbol"],
                )
                rows_added += added
                actions = self.provider.get_actions(download.psx_symbol, start, end)
                if not actions.empty:
                    actions = actions[list(CORPORATE_ACTION_COLUMNS)]
                    _, added_actions = self.store.upsert(
                        actions, self.store.corporate_actions_path,
                        ["symbol", "action_date", "action_type", "source_reference"],
                        ["ratio", "cash_amount"], ["action_date", "symbol"],
                    )
                    actions_added += added_actions
                successful_symbols.append(download.psx_symbol)
                coverage_rows.append(self._coverage_row(download, normalized, actions))
                status_rows.append(self._status_row(download, len(normalized), ""))
            except (RuntimeError, ValueError, OSError, DataConflictError) as exc:
                failure = {"symbol": download.psx_symbol, "yahoo_symbol": download.yahoo_symbol, "error": str(exc)}
                failures.append(failure)
                status = download.status if download.status != "SUCCESS" else "FAILED"
                status_rows.append(self._status_row(download, 0, str(exc), status))
                coverage_rows.append(self._coverage_row(download, pd.DataFrame(), pd.DataFrame(), status))

        index_rows_added = 0
        index_error = ""
        index_download = self._cached_or_download(
            self.settings.benchmark_symbol, start, end, force, "indices"
        )
        if index_download.status == "SUCCESS":
            try:
                raw_reference = self._preserve_raw(index_download, start, end, "indices")
                indices = self.provider.normalize_download(index_download).rename(columns={"symbol": "index"})
                indices["raw_artifact"] = raw_reference
                indices = indices[[column for column in CANONICAL_INDEX_COLUMNS if column in indices]]
                _, index_rows_added = self.store.upsert(
                    indices, self.store.kse100_path, ["index", "date"],
                    ["open", "high", "low", "close", "volume"], ["date", "index"],
                )
            except (ValueError, OSError, DataConflictError) as exc:
                index_error = str(exc)
        else:
            index_error = index_download.error

        equities = self.store.read(self.store.equities_path)
        equities, rejected_rows = self._quarantine_invalid_ohlc(equities)
        rows_added = max(len(equities) - initial_equity_rows, 0)
        indices = self.store.read(self.store.kse100_path)
        metadata = read_stock_metadata(str(symbols_file))
        if not equities.empty:
            DatasetStore._atomic_parquet(build_stock_master(equities, metadata), self.store.stock_master_path)
        actions = self.store.read(self.store.corporate_actions_path)
        discontinuities = reconcile_discontinuities(
            detect_discontinuities(equities, self.settings.extreme_return_threshold), actions
        )
        DatasetStore._atomic_parquet(
            discontinuities,
            self.store.root / "normalized" / "discontinuities.parquet",
        )
        validator = DataValidator(
            self.settings.extreme_return_threshold,
            self.settings.validation.get("expected_session_coverage", 0.90),
            self.settings.minimum_history,
        )
        quality = validator.report(
            validator.validate(equities, indices["date"] if not indices.empty else None),
            self.store.root / "normalized" / "quality-report.json",
        )
        successful = len(successful_symbols)
        status = "FAILED" if successful == 0 else "PARTIAL" if failures or index_error else "SUCCESS"
        latest_sync = {
            "status": status, "timestamp": datetime.now(UTC).isoformat(),
            "provider": "yahoo_finance", "provider_trust": "RESEARCH_SECONDARY",
            "attempted": len(downloads), "successful": successful, "failures": len(failures),
            "benchmark_status": "SUCCESS" if not index_error else "FAILED",
        }
        manifest = generate_manifest(self.store, "yahoo_finance", quality, latest_sync)
        self._write_status_reports(status_rows, coverage_rows, equities, indices)
        self._write_symbol_map(downloads)
        result = YahooBootstrapResult(
            status=status, attempted=len(downloads), successful=successful,
            failed=len(failures), rows_added=rows_added, index_rows_added=index_rows_added,
            action_rows_added=actions_added, rejected_rows=rejected_rows,
            failures=failures, manifest=manifest,
        )
        sync_report = self.store.root / "normalized" / "yahoo-sync-report.json"
        sync_report.write_text(
            json.dumps({**asdict(result), "benchmark_error": index_error, "force": force}, indent=2, default=str),
            encoding="utf-8",
        )
        return result

    def _cached_or_download(
        self, symbol: str, start: date, end: date, force: bool, kind: str = "equities"
    ) -> YahooDownload:
        yahoo_symbol = self.provider.resolve_symbol(symbol)
        cache_path = self._cache_path(yahoo_symbol, start, end, kind)
        if cache_path.exists() and not force:
            frame = pd.read_csv(cache_path, index_col=0, parse_dates=[0])
            return YahooDownload(symbol.upper(), yahoo_symbol, "SUCCESS", frame, 0)
        return self.provider.download_symbol(symbol, start, end)

    def _quarantine_invalid_ohlc(self, equities: pd.DataFrame) -> tuple[pd.DataFrame, int]:
        if equities.empty:
            return equities, 0
        yahoo = equities.get("provider", pd.Series("", index=equities.index)) == "yahoo_finance"
        open_invalid = (equities["open"] < equities["low"]) | (equities["open"] > equities["high"])
        close_invalid = (equities["close"] < equities["low"]) | (equities["close"] > equities["high"])
        mask = yahoo & (open_invalid | close_invalid)
        rejected = equities.loc[mask].copy()
        rejected["rejection_reason"] = ""
        rejected.loc[open_invalid[mask], "rejection_reason"] = "OPEN_OUTSIDE_HIGH_LOW"
        both = open_invalid & close_invalid & mask
        rejected.loc[close_invalid[mask], "rejection_reason"] = "CLOSE_OUTSIDE_HIGH_LOW"
        rejected.loc[both, "rejection_reason"] = "OPEN_AND_CLOSE_OUTSIDE_HIGH_LOW"
        report = Path("reports/task4/yahoo_rejected_rows.csv")
        report.parent.mkdir(parents=True, exist_ok=True)
        rejected.to_csv(report, index=False)
        cleaned = equities.loc[~mask].reset_index(drop=True)
        DatasetStore._atomic_parquet(cleaned, self.store.equities_path)
        return cleaned, len(rejected)

    def _preserve_raw(
        self, download: YahooDownload, start: date, end: date, kind: str = "equities"
    ) -> str:
        cache_path = self._cache_path(download.yahoo_symbol, start, end, kind)
        cache_path.parent.mkdir(parents=True, exist_ok=True)
        download.frame.to_csv(cache_path, index=True)
        artifact = MarketDataArtifact(
            cache_path, kind, f"Yahoo Finance daily history: {download.yahoo_symbol}",
            None, "yahoo",
        )
        return self.raw_store.preserve(artifact)

    def _cache_path(self, yahoo_symbol: str, start: date, end: date, kind: str) -> Path:
        name = f"{yahoo_symbol.replace('^', 'INDEX_')}_{start}_{end}.csv"
        return self.cache_dir / kind / name

    @staticmethod
    def _status_row(
        download: YahooDownload, rows: int, error: str, status: str | None = None
    ) -> dict[str, object]:
        return {
            "psx_symbol": download.psx_symbol, "yahoo_symbol": download.yahoo_symbol,
            "status": status or download.status, "rows": rows,
            "attempts": download.attempts, "error": error,
        }

    @staticmethod
    def _coverage_row(
        download: YahooDownload, frame: pd.DataFrame, actions: pd.DataFrame,
        status: str | None = None,
    ) -> dict[str, object]:
        dividends = int((actions.get("action_type", pd.Series(dtype="string")) == "CASH_DIVIDEND").sum())
        splits = int((actions.get("action_type", pd.Series(dtype="string")) == "STOCK_SPLIT").sum())
        calendar_days = (frame["date"].max() - frame["date"].min()).days if len(frame) > 1 else 0
        expected_weekdays = len(pd.bdate_range(frame["date"].min(), frame["date"].max())) if not frame.empty else 0
        frequency = len(frame) / expected_weekdays if expected_weekdays else 0.0
        return {
            "symbol": download.psx_symbol, "yahoo_symbol": download.yahoo_symbol,
            "first_date": str(frame["date"].min().date()) if not frame.empty else "",
            "last_date": str(frame["date"].max().date()) if not frame.empty else "",
            "rows": len(frame), "trading_frequency": min(frequency, 1.0),
            "missing_ratio": max(0.0, 1.0 - frequency),
            "dividend_events": dividends, "split_events": splits,
            "status": status or download.status, "calendar_days": calendar_days,
        }

    @staticmethod
    def _write_status_reports(
        status_rows: list[dict[str, object]], coverage_rows: list[dict[str, object]],
        equities: pd.DataFrame, indices: pd.DataFrame,
    ) -> None:
        status_path = Path("reports/yahoo_download_status.csv")
        coverage_path = Path("reports/task4/yahoo_coverage.csv")
        status_path.parent.mkdir(parents=True, exist_ok=True)
        coverage_path.parent.mkdir(parents=True, exist_ok=True)
        status = pd.DataFrame(status_rows)
        coverage = pd.DataFrame(coverage_rows)
        if not equities.empty:
            actual = equities.groupby("symbol").agg(
                first_date=("date", "min"), last_date=("date", "max"), rows=("date", "size")
            )
            for index, row in coverage.iterrows():
                symbol = row["symbol"]
                if symbol not in actual.index:
                    continue
                first, last, rows = actual.loc[symbol, ["first_date", "last_date", "rows"]]
                expected = len(pd.bdate_range(first, last))
                frequency = int(rows) / expected if expected else 0.0
                coverage.loc[index, ["first_date", "last_date", "rows", "trading_frequency", "missing_ratio"]] = [
                    str(pd.Timestamp(first).date()), str(pd.Timestamp(last).date()), int(rows),
                    min(frequency, 1.0), max(0.0, 1.0 - frequency),
                ]
                status.loc[status["psx_symbol"] == symbol, "rows"] = int(rows)
        status.to_csv(status_path, index=False)
        coverage.drop(columns=["calendar_days"], errors="ignore").to_csv(coverage_path, index=False)
        benchmark_path = Path("reports/task4/kse100_coverage.json")
        equity_dates = set(pd.to_datetime(equities["date"]).dt.normalize()) if not equities.empty else set()
        index_dates = set(pd.to_datetime(indices["date"]).dt.normalize()) if not indices.empty else set()
        index_end = pd.Timestamp(indices["date"].max()) if not indices.empty else None
        missing_after_end = len([value for value in equity_dates if index_end is not None and value > index_end])
        benchmark = {
            "start": str(indices["date"].min().date()) if not indices.empty else None,
            "end": str(indices["date"].max().date()) if not indices.empty else None,
            "rows": len(indices),
            "missing_equity_sessions_total": len(equity_dates - index_dates),
            "missing_equity_sessions_after_benchmark_end": missing_after_end,
            "stale_calendar_days": (
                int((pd.Timestamp(equities["date"].max()) - index_end).days)
                if not equities.empty and index_end is not None else None
            ),
        }
        benchmark_path.write_text(json.dumps(benchmark, indent=2), encoding="utf-8")

    def _write_symbol_map(self, downloads: list[YahooDownload]) -> None:
        existing = self.provider._symbol_map()
        now = datetime.now(UTC).date().isoformat()
        updates = pd.DataFrame([{
            "psx_symbol": item.psx_symbol, "yahoo_symbol": item.yahoo_symbol,
            "status": item.status, "verified_at": now,
            "notes": item.error[:240] if item.error else "Verified by non-empty daily history",
        } for item in downloads])
        combined = pd.concat([existing, updates], ignore_index=True)
        combined = combined.drop_duplicates("psx_symbol", keep="last").sort_values("psx_symbol")
        self.provider.symbol_map_path.parent.mkdir(parents=True, exist_ok=True)
        combined.to_csv(self.provider.symbol_map_path, index=False)
