"""Tests for token exhaustion guard: balance check, soft stop, auto-resume."""

from __future__ import annotations

from unittest.mock import MagicMock, patch

import pytest

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


def _settings(tmp_path, **overrides) -> Settings:
    defaults = dict(
        dgred_api_key="test",
        dry_run=True,
        db_path=str(tmp_path / "test.db"),
        airminal_billing_key="fake_key",
        airminal_billing_url="https://mock.airminal.com/api",
        billing_mode="live",
        tokens_per_scan=0.3,
        tokens_per_lead_67=2.0,
        tokens_per_lead_68=0.8,
        tokens_per_lead_69=0.0,
        billing_max_tokens_per_cycle=200.0,
        sendgrid_api_key="",
        notify_email_to="",
        notify_email_from="test@test.com",
        fb_cookies="",
        fb_proxy_url="",
        warmup_until="",
        max_daily_scans=30,
    )
    defaults.update(overrides)
    return Settings(**defaults)


def _ledger_with_mock_client(tmp_path, balance: float = 100.0, **overrides):
    """Create a ledger with a mocked AirminalClient returning the given balance."""
    # Create account_guard table (needed for tokens_exhausted flag)
    import sqlite3
    db_path = str(tmp_path / "test.db")
    conn = sqlite3.connect(db_path)
    conn.execute("""
        CREATE TABLE IF NOT EXISTS account_guard (
            key TEXT PRIMARY KEY,
            value TEXT NOT NULL,
            updated_at TEXT NOT NULL
        )
    """)
    conn.commit()
    conn.close()

    s = _settings(tmp_path, db_path=db_path, **overrides)
    ledger = BillingLedger(s, cycle_id=1)

    mock_client = MagicMock()
    mock_client.check_balance.return_value = BalanceResult(
        success=True, balance=balance, locked=0.0,
    )
    mock_client.charge.return_value = ChargeResult(
        success=True, charged=0.3, balance_after=balance - 0.3, error="",
    )
    ledger._client = mock_client
    return ledger, mock_client


# ---------------------------------------------------------------------------
# Pre-flight balance check
# ---------------------------------------------------------------------------

class TestBalancePreFlight:
    def test_sufficient_balance_returns_true(self, tmp_path):
        ledger, _ = _ledger_with_mock_client(tmp_path, balance=100.0)
        assert ledger.check_balance_sufficient(min_tokens=0.3) is True
        assert ledger.is_tokens_exhausted() is False
        ledger.close()

    def test_zero_balance_returns_false(self, tmp_path):
        ledger, _ = _ledger_with_mock_client(tmp_path, balance=0.0)
        assert ledger.check_balance_sufficient(min_tokens=0.3) is False
        assert ledger.is_tokens_exhausted() is True
        ledger.close()

    def test_below_threshold_returns_false(self, tmp_path):
        ledger, _ = _ledger_with_mock_client(tmp_path, balance=0.2)
        assert ledger.check_balance_sufficient(min_tokens=0.3) is False
        assert ledger.is_tokens_exhausted() is True
        ledger.close()

    def test_api_error_fails_open(self, tmp_path):
        """If balance check API fails, proceed cautiously (don't block)."""
        ledger, mock_client = _ledger_with_mock_client(tmp_path)
        mock_client.check_balance.side_effect = RuntimeError("timeout")
        assert ledger.check_balance_sufficient() is True
        ledger.close()


# ---------------------------------------------------------------------------
# Auto-resume when balance restored
# ---------------------------------------------------------------------------

class TestAutoResume:
    def test_restored_balance_clears_flag(self, tmp_path):
        ledger, mock_client = _ledger_with_mock_client(tmp_path, balance=0.0)

        # First check: exhausted
        assert ledger.check_balance_sufficient(min_tokens=0.3) is False
        assert ledger.is_tokens_exhausted() is True

        # Simulate top-up
        mock_client.check_balance.return_value = BalanceResult(
            success=True, balance=50.0, locked=0.0,
        )

        # Second check: restored
        assert ledger.check_balance_sufficient(min_tokens=0.3) is True
        assert ledger.is_tokens_exhausted() is False
        ledger.close()

    def test_flag_persists_across_instances(self, tmp_path):
        import sqlite3
        db_path = str(tmp_path / "test.db")

        # Create tables
        conn = sqlite3.connect(db_path)
        conn.execute("""
            CREATE TABLE IF NOT EXISTS account_guard (
                key TEXT PRIMARY KEY,
                value TEXT NOT NULL,
                updated_at TEXT NOT NULL
            )
        """)
        conn.commit()
        conn.close()

        # Instance 1: exhaust
        l1, _ = _ledger_with_mock_client(tmp_path, balance=0.0)
        l1.check_balance_sufficient(min_tokens=0.3)
        assert l1.is_tokens_exhausted() is True
        l1.close()

        # Instance 2: still exhausted (flag persisted)
        s = _settings(tmp_path, db_path=db_path)
        l2 = BillingLedger(s, cycle_id=2)
        mock2 = MagicMock()
        mock2.check_balance.return_value = BalanceResult(
            success=True, balance=0.1, locked=0.0,
        )
        l2._client = mock2
        assert l2.is_tokens_exhausted() is True
        assert l2.check_balance_sufficient(min_tokens=0.3) is False
        l2.close()

        # Instance 3: restored
        l3 = BillingLedger(s, cycle_id=3)
        mock3 = MagicMock()
        mock3.check_balance.return_value = BalanceResult(
            success=True, balance=100.0, locked=0.0,
        )
        l3._client = mock3
        assert l3.check_balance_sufficient(min_tokens=0.3) is True
        assert l3.is_tokens_exhausted() is False
        l3.close()


# ---------------------------------------------------------------------------
# Mid-cycle InsufficientTokensError
# ---------------------------------------------------------------------------

class TestMidCycleExhaustion:
    def test_insufficient_tokens_sets_flag_and_halts(self, tmp_path):
        ledger, mock_client = _ledger_with_mock_client(tmp_path, balance=100.0)
        mock_client.charge.side_effect = InsufficientTokensError("no funds")

        result = ledger.charge_scan("g1", "Group 1")
        assert result is False
        assert ledger.halted is True
        assert ledger.is_tokens_exhausted() is True
        ledger.close()

    def test_successful_charge_does_not_set_flag(self, tmp_path):
        ledger, _ = _ledger_with_mock_client(tmp_path, balance=100.0)
        result = ledger.charge_scan("g1", "Group 1")
        assert result is True
        assert ledger.is_tokens_exhausted() is False
        ledger.close()


# ---------------------------------------------------------------------------
# Shadow mode: no balance checks
# ---------------------------------------------------------------------------

class TestShadowMode:
    def test_shadow_check_balance_always_true(self, tmp_path):
        s = _settings(tmp_path, billing_mode="shadow", airminal_billing_key="")
        ledger = BillingLedger(s, cycle_id=1)
        # No client in shadow mode
        assert ledger.check_balance_sufficient() is True
        assert ledger.is_tokens_exhausted() is False
        ledger.close()


# ---------------------------------------------------------------------------
# Email alarm (mocked)
# ---------------------------------------------------------------------------

class TestTokenAlarmEmail:
    @patch("sendgrid.SendGridAPIClient")
    def test_exhaustion_alarm_sends(self, MockSG, tmp_path):
        s = _settings(tmp_path, sendgrid_api_key="SG.fake", notify_email_to="a@b.com")
        ledger = BillingLedger(s, cycle_id=1)
        ledger.send_tokens_alarm(s, exhausted=True, balance=0.1)
        MockSG.assert_called()
        ledger.close()

    @patch("sendgrid.SendGridAPIClient")
    def test_restoration_alarm_sends(self, MockSG, tmp_path):
        s = _settings(tmp_path, sendgrid_api_key="SG.fake", notify_email_to="a@b.com")
        ledger = BillingLedger(s, cycle_id=1)
        ledger.send_tokens_alarm(s, exhausted=False, balance=50.0)
        MockSG.assert_called()
        ledger.close()
