"""Billing ledger: shadow/live token charging with dedup, circuit breaker, and logging.

In SHADOW mode: calculates would_charge, logs it, but never calls the airminal API.
In LIVE mode: checks balance, charges via API BEFORE the action, blocks on failure.
"""

from __future__ import annotations

import logging
import sqlite3
from datetime import datetime
from pathlib import Path

from agent_samochodowy.billing.client import (
    AirminalClient,
    ChargeResult,
    InsufficientTokensError,
)
from agent_samochodowy.config import Settings

logger = logging.getLogger(__name__)

# Status name → settings field mapping
_LEAD_STATUS_TOKEN_FIELD = {
    "Dopasowano": "tokens_per_lead_67",
    "Do weryfikacji": "tokens_per_lead_68",
    "Brak dopasowania": "tokens_per_lead_69",
}


class BillingLedger:
    """Manages token charges per cycle with shadow/live modes."""

    def __init__(self, settings: Settings, cycle_id: int | str = 0) -> None:
        self._settings = settings
        self._mode = settings.billing_mode  # "shadow" or "live"
        self._cycle_id = cycle_id
        self._max_per_cycle = settings.billing_max_tokens_per_cycle
        self._cycle_total = 0.0
        self._halted = False

        # Airminal API client — only instantiated if key is present
        self._client: AirminalClient | None = None
        if settings.airminal_billing_key:
            self._client = AirminalClient(
                settings.airminal_billing_url,
                settings.airminal_billing_key,
            )

        # SQLite ledger for dedup + history
        db_path = settings.db_path
        Path(db_path).parent.mkdir(parents=True, exist_ok=True)
        self._conn = sqlite3.connect(db_path)
        self._conn.row_factory = sqlite3.Row
        self._init_schema()

    @property
    def is_live(self) -> bool:
        return self._mode == "live"

    @property
    def halted(self) -> bool:
        return self._halted

    @property
    def cycle_total(self) -> float:
        return self._cycle_total

    # ------------------------------------------------------------------
    # Public API
    # ------------------------------------------------------------------

    def charge_scan(self, group_id: str, group_name: str) -> bool:
        """Charge for a group scan. Returns True if action should proceed."""
        tokens = self._settings.tokens_per_scan
        if tokens <= 0:
            return True
        reference = f"scan-{group_id}-{self._cycle_id}"
        task = f"skan {group_name[:50]}"
        return self._charge_or_log(tokens, task, reference, "scan")

    def charge_lead(
        self, status_name: str, client_id: int | str, post_id: str,
    ) -> bool:
        """Charge for a lead write. Returns True if action should proceed."""
        field = _LEAD_STATUS_TOKEN_FIELD.get(status_name)
        if not field:
            return True
        tokens = getattr(self._settings, field, 0.0)
        if tokens <= 0:
            return True
        reference = f"lead-{client_id}-{post_id}"
        task = f"lead {status_name}"
        return self._charge_or_log(tokens, task, reference, "lead")

    def invalidate_empty_scans(self, db_path: str) -> dict:
        """Mark billing entries as invalid for cycles where posts_fetched=0.

        Returns summary dict with counts and tokens invalidated.
        """
        conn = sqlite3.connect(db_path)
        # Find cycle IDs with 0 posts
        empty_cycles = conn.execute(
            "SELECT id FROM cycle_log WHERE posts_fetched = 0"
        ).fetchall()
        empty_ids = [str(r[0]) for r in empty_cycles]
        if not empty_ids:
            conn.close()
            return {"invalidated": 0, "tokens": 0.0}

        placeholders = ",".join("?" for _ in empty_ids)
        # Count before update
        cur = conn.execute(
            f"SELECT COUNT(*), COALESCE(SUM(tokens), 0) FROM billing_log "
            f"WHERE cycle_id IN ({placeholders}) AND action_type = 'scan' "
            f"AND result != 'invalid_empty_scan'",
            empty_ids,
        )
        count, tokens = cur.fetchone()

        # Mark invalid
        conn.execute(
            f"UPDATE billing_log SET result = 'invalid_empty_scan' "
            f"WHERE cycle_id IN ({placeholders}) AND action_type = 'scan' "
            f"AND result != 'invalid_empty_scan'",
            empty_ids,
        )
        conn.commit()
        conn.close()

        logger.info(
            "Invalidated %d scan charges (%.1f tokens) from %d empty cycles",
            count, tokens, len(empty_ids),
        )
        return {"invalidated": count, "tokens": float(tokens), "empty_cycles": len(empty_ids)}

    # ------------------------------------------------------------------
    # Token exhaustion guard
    # ------------------------------------------------------------------

    def check_balance_sufficient(self, min_tokens: float = 0.3) -> bool:
        """Check if org has enough tokens to proceed. Returns True if OK.

        Sets tokens_exhausted flag in DB if insufficient.
        Auto-clears flag if balance is restored (soft stop).
        """
        if not self._client:
            return True  # no client = shadow mode, always OK

        try:
            result = self._client.check_balance()
            available = result.balance - result.locked
        except Exception:
            logger.exception("Balance check failed — proceeding cautiously")
            return True  # fail-open: don't block on API errors

        if available < min_tokens:
            if not self._is_tokens_exhausted():
                self._set_tokens_exhausted(True, available)
                logger.error(
                    "TOKENS EXHAUSTED: available=%.1f < min=%.1f — halting",
                    available, min_tokens,
                )
            return False

        # Balance restored — auto-clear if previously exhausted
        if self._is_tokens_exhausted():
            self._set_tokens_exhausted(False, available)
            logger.info(
                "TOKENS RESTORED: available=%.1f — resuming",
                available,
            )
        return True

    def is_tokens_exhausted(self) -> bool:
        return self._is_tokens_exhausted()

    def _is_tokens_exhausted(self) -> bool:
        cur = self._conn.execute(
            "SELECT 1 FROM account_guard WHERE key = 'tokens_exhausted' AND value = 'true'"
        )
        return cur.fetchone() is not None

    def _set_tokens_exhausted(self, exhausted: bool, balance: float) -> None:
        from datetime import datetime as _dt
        now = _dt.now().isoformat()
        if exhausted:
            self._conn.execute(
                "INSERT OR REPLACE INTO account_guard (key, value, updated_at) "
                "VALUES ('tokens_exhausted', 'true', ?)", (now,)
            )
            self._conn.execute(
                "INSERT OR REPLACE INTO account_guard (key, value, updated_at) "
                "VALUES ('tokens_exhausted_balance', ?, ?)", (str(balance), now)
            )
        else:
            self._conn.execute(
                "DELETE FROM account_guard WHERE key IN "
                "('tokens_exhausted', 'tokens_exhausted_balance')"
            )
        self._conn.commit()

    def send_tokens_alarm(self, settings, exhausted: bool, balance: float) -> None:
        """Send email about token exhaustion or restoration."""
        if not settings.sendgrid_api_key or not settings.notify_email_to:
            return
        try:
            from sendgrid import SendGridAPIClient
            from sendgrid.helpers.mail import Mail

            recipients = [e.strip() for e in settings.notify_email_to.split(",") if e.strip()]
            if exhausted:
                subject = "ALARM: Agent Samochodowy — tokeny wyczerpane"
                html = (
                    f"<h2 style='color:#dc2626'>Tokeny na koncie wyczerpane</h2>"
                    f"<p>Saldo: <strong>{balance:.1f}</strong> tokenów</p>"
                    f"<p>Agent został <strong>automatycznie wstrzymany</strong>.</p>"
                    f"<p>Doładuj tokeny, aby wznowić skanowanie.</p>"
                )
            else:
                subject = "INFO: Agent Samochodowy — wznowiono (tokeny doładowane)"
                html = (
                    f"<h2 style='color:#16a34a'>Saldo doładowane — agent wznowiony</h2>"
                    f"<p>Saldo: <strong>{balance:.1f}</strong> tokenów</p>"
                    f"<p>Skanowanie zostało automatycznie wznowione.</p>"
                )

            for recipient in recipients:
                sg = SendGridAPIClient(settings.sendgrid_api_key)
                sg.send(Mail(
                    from_email=settings.notify_email_from,
                    to_emails=recipient,
                    subject=subject,
                    html_content=html,
                ))
            logger.info("Token %s alarm sent to %d recipient(s)",
                        "exhausted" if exhausted else "restored", len(recipients))
        except Exception:
            logger.exception("Failed to send token alarm email")

    def close(self) -> None:
        self._conn.commit()
        self._conn.close()
        if self._client:
            self._client.close()

    # ------------------------------------------------------------------
    # Schema
    # ------------------------------------------------------------------

    def _init_schema(self) -> None:
        self._conn.execute("""
            CREATE TABLE IF NOT EXISTS billing_log (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                timestamp TEXT NOT NULL,
                cycle_id TEXT,
                action_type TEXT NOT NULL,
                task TEXT NOT NULL,
                reference TEXT NOT NULL,
                tokens REAL NOT NULL,
                mode TEXT NOT NULL,
                result TEXT NOT NULL,
                balance_after REAL
            )
        """)
        self._conn.execute("""
            CREATE TABLE IF NOT EXISTS charged_references (
                reference TEXT PRIMARY KEY,
                charged_at TEXT NOT NULL
            )
        """)
        self._conn.execute("""
            CREATE TABLE IF NOT EXISTS account_guard (
                key TEXT PRIMARY KEY,
                value TEXT NOT NULL,
                updated_at TEXT NOT NULL
            )
        """)
        self._conn.commit()

    # ------------------------------------------------------------------
    # Core logic
    # ------------------------------------------------------------------

    def _charge_or_log(
        self, tokens: float, task: str, reference: str, action_type: str,
    ) -> bool:
        """Unified charge/log entry point. Returns True if action should proceed."""
        # Circuit breaker
        if self._halted:
            logger.warning("Billing HALTED — skipping %s (%s)", reference, task)
            return self._mode != "live"  # shadow: proceed anyway

        if self._cycle_total + tokens > self._max_per_cycle:
            self._halted = True
            logger.error(
                "BILLING CIRCUIT BREAKER: cycle total %.1f + %.1f > %.1f — HALT",
                self._cycle_total, tokens, self._max_per_cycle,
            )
            self._log_entry(action_type, task, reference, tokens, "HALTED", None)
            return self._mode != "live"

        # Dedup: check if this reference was already charged
        if self._is_duplicate(reference):
            logger.debug("Billing dedup: %s already charged, skipping", reference)
            return True  # action already paid for

        if self.is_live:
            return self._charge_live(tokens, task, reference, action_type)
        return self._charge_shadow(tokens, task, reference, action_type)

    def _charge_shadow(
        self, tokens: float, task: str, reference: str, action_type: str,
    ) -> bool:
        """Shadow mode: log would_charge, always proceed."""
        self._cycle_total += tokens
        self._mark_charged(reference)
        self._log_entry(action_type, task, reference, tokens, "would_charge", None)
        logger.info(
            "[SHADOW] would_charge %.1f tokens for %s (ref=%s, cycle_total=%.1f)",
            tokens, task, reference, self._cycle_total,
        )
        return True

    def _charge_live(
        self, tokens: float, task: str, reference: str, action_type: str,
    ) -> bool:
        """Live mode: charge via API BEFORE action. Block on failure."""
        if not self._client:
            logger.error("LIVE mode but no airminal API key — blocking action %s", reference)
            self._log_entry(action_type, task, reference, tokens, "error_no_key", None)
            return False

        try:
            result = self._client.charge(tokens, task, reference)
            self._cycle_total += tokens
            self._mark_charged(reference)
            self._log_entry(
                action_type, task, reference, tokens,
                "charged", result.balance_after,
            )
            logger.info(
                "[LIVE] charged %.1f tokens for %s (ref=%s, balance=%.1f)",
                tokens, task, reference, result.balance_after,
            )
            return True
        except InsufficientTokensError:
            self._log_entry(action_type, task, reference, tokens, "insufficient", None)
            self._set_tokens_exhausted(True, 0)
            self._halted = True
            logger.error(
                "[LIVE] INSUFFICIENT TOKENS for %s (%.1f needed, ref=%s) — "
                "setting tokens_exhausted, halting cycle",
                task, tokens, reference,
            )
            return False
        except Exception:
            self._log_entry(action_type, task, reference, tokens, "error", None)
            logger.exception(
                "[LIVE] airminal API error for %s — fail-safe: blocking action", reference,
            )
            return False

    # ------------------------------------------------------------------
    # Dedup
    # ------------------------------------------------------------------

    def _is_duplicate(self, reference: str) -> bool:
        cur = self._conn.execute(
            "SELECT 1 FROM charged_references WHERE reference = ?", (reference,)
        )
        return cur.fetchone() is not None

    def _mark_charged(self, reference: str) -> None:
        self._conn.execute(
            "INSERT OR IGNORE INTO charged_references (reference, charged_at) VALUES (?, ?)",
            (reference, datetime.now().isoformat()),
        )
        self._conn.commit()

    # ------------------------------------------------------------------
    # Logging
    # ------------------------------------------------------------------

    def _log_entry(
        self,
        action_type: str,
        task: str,
        reference: str,
        tokens: float,
        result: str,
        balance_after: float | None,
    ) -> None:
        self._conn.execute("""
            INSERT INTO billing_log
                (timestamp, cycle_id, action_type, task, reference, tokens, mode, result, balance_after)
            VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
        """, (
            datetime.now().isoformat(),
            str(self._cycle_id),
            action_type,
            task,
            reference,
            tokens,
            self._mode,
            result,
            balance_after,
        ))
        self._conn.commit()
