from __future__ import annotations

from dataclasses import dataclass
from decimal import Decimal
from sqlite3 import Connection
from typing import Protocol, Sequence

from jarvis_finance.audit.log import record_audit_event
from jarvis_finance.imports.common import stable_id, utc_now
from jarvis_finance.quality.alerts import create_alert


class FxProvider(Protocol):
    def get_rate(self, base_currency: str, quote_currency: str, rate_date: str | None = None) -> Decimal | None: ...


@dataclass
class FxUpdateResult:
    updated_count: int = 0
    cached_count: int = 0
    skipped_count: int = 0
    warning_count: int = 0
    error_count: int = 0
    resolved_alerts: int = 0
    dry_run: bool = False


@dataclass
class FxRecheckResult:
    transaction_count: int = 0
    chf_positions: int = 0
    foreign_currency_positions: int = 0
    justified_missing_fx_alerts: int = 0
    false_positive_missing_fx_alerts: int = 0
    corrected_transactions: int = 0
    resolved_alerts: int = 0
    audit_events: int = 0
    dry_run: bool = False


@dataclass(frozen=True)
class FxResolutionResult:
    base_currency: str
    quote_currency: str = "CHF"
    rate: Decimal | None = None
    source: str = "missing"
    status: str = "missing"
    rate_date: str | None = None
    warning: str | None = None


def _resolve_alerts(conn: Connection, *, rule_id: str, entity_ids: Sequence[str] | None = None) -> int:
    now = utc_now()
    params: list[object] = [rule_id]
    where = "rule_id=? AND status='active'"
    if entity_ids is not None:
        if not entity_ids:
            return 0
        where += " AND entity_id IN (" + ",".join("?" for _ in entity_ids) + ")"
        params.extend(entity_ids)
    rows = conn.execute(f"SELECT alert_id FROM alerts WHERE {where}", tuple(params)).fetchall()
    for row in rows:
        conn.execute("UPDATE alerts SET status='resolved', resolved_at=?, last_seen_at=? WHERE alert_id=?", (now, now, row["alert_id"]))
    return len(rows)


def _transaction_instruments_for_currency(conn: Connection, currency: str) -> list[str]:
    return [
        row["instrument_id"]
        for row in conn.execute(
            """
            SELECT DISTINCT instrument_id FROM transactions
            WHERE instrument_id IS NOT NULL AND currency_original=? AND COALESCE(is_voided,0)=0
            """,
            (currency.upper(),),
        ).fetchall()
    ]


def _mark_chf_transactions_not_needed(conn: Connection, *, dry_run: bool = False) -> int:
    now = utc_now()
    rows = conn.execute(
        """
        SELECT transaction_id, fx_rate_to_chf, fx_source, fx_status FROM transactions
        WHERE currency_original='CHF'
          AND (fx_status IS NULL OR fx_status='missing' OR fx_rate_to_chf IS NULL OR fx_source IS NULL)
        """
    ).fetchall()
    if dry_run:
        return len(rows)
    for row in rows:
        conn.execute(
            "UPDATE transactions SET fx_rate_to_chf='1', fx_source='not_needed', fx_status='not_needed', updated_at=? WHERE transaction_id=?",
            (now, row["transaction_id"]),
        )
        record_audit_event(
            conn,
            source="fx_recheck",
            action="set_chf_fx_not_needed",
            entity_type="transaction",
            entity_id=row["transaction_id"],
            old_values={"fx_rate_to_chf": row["fx_rate_to_chf"], "fx_source": row["fx_source"], "fx_status": row["fx_status"]},
            new_values={"fx_rate_to_chf": "1", "fx_source": "not_needed", "fx_status": "not_needed"},
            user_text_note="CHF original currency requires no FX conversion.",
        )
    if rows:
        _resolve_alerts(conn, rule_id="missing_fx", entity_ids=_transaction_instruments_for_currency(conn, "CHF"))
        conn.commit()
    return len(rows)


def recheck_transaction_fx_status(conn: Connection, *, source_type: str | None = None, dry_run: bool = False) -> FxRecheckResult:
    result = FxRecheckResult(dry_run=dry_run)
    params: list[object] = []
    where = "COALESCE(is_voided,0)=0 AND transaction_type='initial_position_snapshot'"
    if source_type:
        where += " AND source_type=?"
        params.append(source_type)
    rows = conn.execute(
        f"""
        SELECT transaction_id, instrument_id, currency_original, fx_status, fx_rate_to_chf, fx_source
        FROM transactions
        WHERE {where}
        """,
        tuple(params),
    ).fetchall()
    result.transaction_count = len(rows)
    chf_ids: list[str] = []
    foreign_ids: list[str] = []
    false_positive_ids: list[str] = []
    justified_ids: list[str] = []
    for row in rows:
        cur = (row["currency_original"] or "").upper()
        inst = row["instrument_id"]
        if cur == "CHF":
            result.chf_positions += 1
            if inst:
                chf_ids.append(inst)
            if row["fx_status"] == "missing" or row["fx_rate_to_chf"] in (None, "") or row["fx_source"] in (None, ""):
                result.corrected_transactions += 1
                if inst:
                    false_positive_ids.append(inst)
                if not dry_run:
                    conn.execute(
                        "UPDATE transactions SET fx_rate_to_chf='1', fx_source='not_needed', fx_status='not_needed', updated_at=? WHERE transaction_id=?",
                        (utc_now(), row["transaction_id"]),
                    )
                    record_audit_event(
                        conn,
                        source="fx_recheck",
                        action="set_chf_fx_not_needed",
                        entity_type="transaction",
                        entity_id=row["transaction_id"],
                        old_values={"fx_rate_to_chf": row["fx_rate_to_chf"], "fx_source": row["fx_source"], "fx_status": row["fx_status"]},
                        new_values={"fx_rate_to_chf": "1", "fx_source": "not_needed", "fx_status": "not_needed"},
                        user_text_note="CHF initial snapshot requires no FX conversion.",
                    )
                    result.audit_events += 1
        else:
            result.foreign_currency_positions += 1
            if inst:
                foreign_ids.append(inst)
            if row["fx_status"] == "missing" or not latest_fx_rate(conn, base_currency=cur, quote_currency="CHF", rate_date=None):
                if inst:
                    justified_ids.append(inst)
    if false_positive_ids:
        result.false_positive_missing_fx_alerts = len(conn.execute(
            "SELECT alert_id FROM alerts WHERE status='active' AND rule_id='missing_fx' AND entity_id IN (" + ",".join("?" for _ in false_positive_ids) + ")",
            tuple(false_positive_ids),
        ).fetchall())
        if not dry_run:
            result.resolved_alerts += _resolve_alerts(conn, rule_id="missing_fx", entity_ids=false_positive_ids)
    if justified_ids:
        result.justified_missing_fx_alerts = len(conn.execute(
            "SELECT alert_id FROM alerts WHERE status='active' AND rule_id='missing_fx' AND entity_id IN (" + ",".join("?" for _ in justified_ids) + ")",
            tuple(justified_ids),
        ).fetchall())
    if not dry_run:
        conn.commit()
    return result


def upsert_fx_rate(
    conn: Connection,
    *,
    base_currency: str,
    quote_currency: str,
    rate_date: str,
    rate: Decimal,
    provider: str,
    rate_type: str,
    quality_status: str = "fresh",
    fetched_at: str | None = None,
    run_id: str | None = None,
) -> str:
    base = base_currency.upper()
    quote = quote_currency.upper()
    existing = conn.execute(
        "SELECT fx_rate_id,rate,quality_status FROM fx_rates WHERE base_currency=? AND quote_currency=? AND rate_date=? AND provider=? AND rate_type=?",
        (base, quote, rate_date, provider, rate_type),
    ).fetchone()
    fx_rate_id = stable_id("fx", base, quote, rate_date, provider, rate_type)
    conn.execute(
        """
        INSERT INTO fx_rates(
            fx_rate_id, base_currency, quote_currency, rate_date, rate, provider,
            rate_type, quality_status, created_at, fetched_at, run_id
        ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
        ON CONFLICT(base_currency, quote_currency, rate_date, provider, rate_type)
        DO UPDATE SET rate=excluded.rate, quality_status=excluded.quality_status,
                      fetched_at=excluded.fetched_at, run_id=excluded.run_id
        """,
        (fx_rate_id, base, quote, rate_date, format(rate, "f"), provider, rate_type, quality_status, utc_now(), fetched_at or utc_now(), run_id),
    )
    if existing and str(existing["rate"]) != format(rate, "f"):
        record_audit_event(
            conn,
            source="daily_market_fx_v1",
            action="fx_rate_provider_correction",
            entity_type="fx_rate",
            entity_id=str(existing["fx_rate_id"]),
            old_values={"rate": existing["rate"], "quality_status": existing["quality_status"]},
            new_values={"rate": format(rate, "f"), "quality_status": quality_status, "run_id": run_id},
            created_by="system",
        )
    return fx_rate_id


def upsert_fx_unavailable(
    conn: Connection,
    *,
    base_currency: str,
    quote_currency: str,
    rate_date: str,
    provider: str,
    error_status: str = "manual_override_required",
    error_message: str | None = None,
) -> str:
    base = base_currency.upper()
    quote = quote_currency.upper()
    fx_rate_id = stable_id("fxerr", base, quote, rate_date, provider, error_status)
    conn.execute(
        """
        INSERT INTO fx_rates(fx_rate_id, base_currency, quote_currency, rate_date, rate, provider, rate_type, quality_status, created_at, error_status, error_message)
        VALUES (?, ?, ?, ?, '', ?, 'close', ?, ?, ?, ?)
        ON CONFLICT(base_currency, quote_currency, rate_date, provider, rate_type)
        DO UPDATE SET quality_status=excluded.quality_status, error_status=excluded.error_status, error_message=excluded.error_message
        """,
        (fx_rate_id, base, quote, rate_date, provider, error_status, utc_now(), error_status, error_message),
    )
    return fx_rate_id


def latest_fx_rate(conn: Connection, *, base_currency: str, quote_currency: str = "CHF", rate_date: str | None = None):
    base = base_currency.upper()
    quote = quote_currency.upper()
    if base == quote:
        return None
    if rate_date:
        return conn.execute(
            """
            SELECT * FROM fx_rates WHERE base_currency=? AND quote_currency=? AND rate_date<=? AND quality_status IN ('fresh','ok','manual')
            ORDER BY rate_date DESC, created_at DESC LIMIT 1
            """,
            (base, quote, rate_date),
        ).fetchone()
    return conn.execute(
        """
        SELECT * FROM fx_rates WHERE base_currency=? AND quote_currency=? AND quality_status IN ('fresh','ok','manual')
        ORDER BY rate_date DESC, created_at DESC LIMIT 1
        """,
        (base, quote),
    ).fetchone()


def get_fx_rate_to_chf(conn: Connection, *, base_currency: str, rate_date: str | None = None, resolve_fixed: bool = False) -> Decimal | None:
    base = base_currency.upper()
    if base == "CHF":
        _mark_chf_transactions_not_needed(conn)
        return Decimal("1")
    row = latest_fx_rate(conn, base_currency=base, quote_currency="CHF", rate_date=rate_date)
    if row is None:
        return None
    if resolve_fixed:
        _resolve_alerts(conn, rule_id="missing_fx", entity_ids=_transaction_instruments_for_currency(conn, base))
        conn.commit()
    return Decimal(str(row["rate"]))


def resolve_fx_rate_to_chf(
    conn: Connection,
    *,
    base_currency: str,
    rate_date: str | None = None,
    providers: Sequence[FxProvider] | None = None,
    persist: bool = True,
    resolve_fixed: bool = True,
) -> FxResolutionResult:
    """Resolve FX to CHF for explicit user actions: Frankfurter primary, then cache/fallback providers.

    Dashboard rendering must not call this. It is used by save/refresh buttons and CLI-style flows.
    """
    base = base_currency.upper()
    effective_date = rate_date or utc_now()[:10]
    if base == "CHF":
        return FxResolutionResult(base_currency=base, rate=Decimal("1"), source="not_needed", status="not_needed", rate_date=effective_date)

    if providers is None:
        from jarvis_finance.fx.providers import FrankfurterFxProvider, TwelveDataFxProvider
        providers = [FrankfurterFxProvider(), TwelveDataFxProvider()]
    provider_list = list(providers)
    primary_providers = [p for p in provider_list if getattr(p, "name", "") == "frankfurter"]
    fallback_providers = [p for p in provider_list if getattr(p, "name", "") != "frankfurter"]
    last_warning: str | None = None

    def try_provider(provider: FxProvider) -> FxResolutionResult | None:
        nonlocal last_warning
        provider_name = getattr(provider, "name", "provider")
        try:
            rate = provider.get_rate(base, "CHF", rate_date)
        except Exception as exc:
            warning = str(exc) or f"{provider_name}_provider_error"
            if provider_name == "frankfurter" or last_warning is None:
                last_warning = warning
            return None
        if rate is None or rate <= 0:
            warning = f"{provider_name}_provider_error"
            if provider_name == "frankfurter" or last_warning is None:
                last_warning = warning
            return None
        if persist:
            upsert_fx_rate(
                conn,
                base_currency=base,
                quote_currency="CHF",
                rate_date=effective_date,
                rate=rate,
                provider=provider_name,
                rate_type="close",
                quality_status="fresh",
            )
            if resolve_fixed:
                _resolve_alerts(conn, rule_id="missing_fx", entity_ids=_transaction_instruments_for_currency(conn, base))
            conn.commit()
        return FxResolutionResult(base_currency=base, rate=rate, source=provider_name, status="ok", rate_date=effective_date)

    for provider in primary_providers:
        result = try_provider(provider)
        if result is not None:
            return result

    row = latest_fx_rate(conn, base_currency=base, quote_currency="CHF", rate_date=effective_date)
    if row is not None and row["rate"] not in {None, ""}:
        if resolve_fixed:
            _resolve_alerts(conn, rule_id="missing_fx", entity_ids=_transaction_instruments_for_currency(conn, base))
            conn.commit()
        return FxResolutionResult(base_currency=base, rate=Decimal(str(row["rate"])), source=f"cache:{row['provider']}", status="ok", rate_date=row["rate_date"])

    for provider in fallback_providers:
        result = try_provider(provider)
        if result is not None:
            return result
    return FxResolutionResult(base_currency=base, rate=None, source="missing", status="missing", rate_date=effective_date, warning=last_warning)

def update_fx_rates(
    conn: Connection,
    *,
    provider: FxProvider,
    currencies: Sequence[str],
    rate_date: str,
    latest: bool = False,
    missing_only: bool = False,
    dry_run: bool = False,
    resolve_fixed: bool = True,
) -> FxUpdateResult:
    result = FxUpdateResult(dry_run=dry_run)
    for cur in dict.fromkeys(c.upper() for c in currencies):
        if cur == "CHF":
            if not dry_run:
                _mark_chf_transactions_not_needed(conn)
            result.cached_count += 1
            continue
        if missing_only and latest_fx_rate(conn, base_currency=cur, quote_currency="CHF", rate_date=rate_date):
            result.cached_count += 1
            continue
        try:
            rate = provider.get_rate(cur, "CHF", None if latest else rate_date)
        except Exception as exc:
            status = getattr(exc, "status", "manual_override_required")
            if status in {"manual_override_required", "paywall", "rate_limited", "forbidden"}:
                result.warning_count += 1
                if not dry_run:
                    upsert_fx_unavailable(conn, base_currency=cur, quote_currency="CHF", rate_date=rate_date, provider=getattr(provider, "name", "unknown"), error_status="manual_override_required", error_message=str(exc))
                    for inst_id in _transaction_instruments_for_currency(conn, cur):
                        create_alert(conn, priority="kritisch", category="fx", entity_type="instrument", entity_id=inst_id, rule_id="historical_fx_unavailable", message="Historical FX is unavailable from provider; manual override is required.", evidence={"base_currency": cur, "quote_currency": "CHF", "rate_date": rate_date}, fingerprint=f"historical_fx_unavailable:{cur}:CHF:{rate_date}")
                        create_alert(conn, priority="kritisch", category="fx", entity_type="instrument", entity_id=inst_id, rule_id="manual_override_required", message="Manual FX override is required; no latest FX fallback is used.", evidence={"base_currency": cur, "quote_currency": "CHF", "rate_date": rate_date}, fingerprint=f"manual_fx_required:{cur}:CHF:{rate_date}")
                continue
            result.error_count += 1
            continue
        if rate is None:
            result.warning_count += 1
            create_alert(conn, priority="kritisch", category="fx", entity_type="currency_pair", entity_id=f"{cur}/CHF", rule_id="missing_fx", message="FX rate to CHF is missing.", evidence={"base_currency": cur, "quote_currency": "CHF", "rate_date": rate_date}, fingerprint=f"missing_fx_rate:{cur}:CHF:{rate_date}")
            continue
        if not dry_run:
            upsert_fx_rate(conn, base_currency=cur, quote_currency="CHF", rate_date=rate_date, rate=rate, provider=getattr(provider, "name", "mock"), rate_type="latest" if latest else "close", quality_status="fresh")
            if resolve_fixed:
                result.resolved_alerts += _resolve_alerts(conn, rule_id="missing_fx", entity_ids=_transaction_instruments_for_currency(conn, cur))
        result.updated_count += 1
    if not dry_run:
        conn.commit()
    return result
