"""Tests for billing: client, ledger, dedup, circuit breaker. Airminal API fully mocked."""

from __future__ import annotations

import logging
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


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def _settings(**overrides) -> Settings:
    defaults = dict(
        dgred_api_key="test",
        dry_run=True,
        airminal_billing_key="fake_key_for_tests",
        airminal_billing_url="https://mock.airminal.com/api",
        billing_mode="shadow",
        tokens_per_scan=0.5,
        tokens_per_lead_67=5.0,
        tokens_per_lead_68=2.0,
        tokens_per_lead_69=0.0,
        billing_max_tokens_per_cycle=200.0,
    )
    defaults.update(overrides)
    return Settings(**defaults)


def _ledger(tmp_path, **overrides) -> BillingLedger:
    s = _settings(db_path=str(tmp_path / "test.db"), **overrides)
    return BillingLedger(s, cycle_id=1)


# ---------------------------------------------------------------------------
# AirminalClient tests (HTTP mocked)
# ---------------------------------------------------------------------------

class TestAirminalClient:
    @patch("agent_samochodowy.billing.client.httpx.Client")
    def test_check_balance(self, MockHttpx):
        mock_http = MagicMock()
        mock_resp = MagicMock()
        mock_resp.json.return_value = {"success": True, "balance": 150.5, "locked": 10.0}
        mock_http.get.return_value = mock_resp
        MockHttpx.return_value = mock_http

        client = AirminalClient("https://api.test/billing", "key123")
        result = client.check_balance()

        assert result.success is True
        assert result.balance == 150.5
        assert result.locked == 10.0
        client.close()

    @patch("agent_samochodowy.billing.client.httpx.Client")
    def test_charge_success(self, MockHttpx):
        mock_http = MagicMock()
        mock_resp = MagicMock()
        mock_resp.status_code = 200
        mock_resp.json.return_value = {
            "success": True, "charged": 5.0, "balance_after": 145.5,
        }
        mock_http.post.return_value = mock_resp
        MockHttpx.return_value = mock_http

        client = AirminalClient("https://api.test/billing", "key123")
        result = client.charge(5.0, "lead Dopasowano", "lead-123-abc")

        assert result.success is True
        assert result.charged == 5.0
        assert result.balance_after == 145.5
        # Verify payload
        call_kwargs = mock_http.post.call_args
        payload = call_kwargs.kwargs["json"]
        assert payload["tokens"] == 5.0
        assert payload["reference"] == "lead-123-abc"
        assert payload["service_type"] == "agent_fb"
        client.close()

    @patch("agent_samochodowy.billing.client.httpx.Client")
    def test_charge_insufficient(self, MockHttpx):
        mock_http = MagicMock()
        mock_resp = MagicMock()
        mock_resp.status_code = 402
        mock_resp.headers = {"content-type": "application/json"}
        mock_resp.json.return_value = {"error": "INSUFFICIENT_TOKENS"}
        mock_http.post.return_value = mock_resp
        MockHttpx.return_value = mock_http

        client = AirminalClient("https://api.test/billing", "key123")
        with pytest.raises(InsufficientTokensError):
            client.charge(5.0, "task", "ref")
        client.close()

    @patch("agent_samochodowy.billing.client.httpx.Client")
    def test_bearer_auth_header(self, MockHttpx):
        mock_http = MagicMock()
        mock_resp = MagicMock()
        mock_resp.json.return_value = {"success": True, "balance": 100, "locked": 0}
        mock_http.get.return_value = mock_resp
        MockHttpx.return_value = mock_http

        client = AirminalClient("https://api.test", "secret_key")
        client.check_balance()

        params = mock_http.get.call_args.kwargs["params"]
        assert params["api_key"] == "secret_key"
        client.close()


# ---------------------------------------------------------------------------
# BillingLedger — shadow mode
# ---------------------------------------------------------------------------

class TestLedgerShadow:
    def test_shadow_always_proceeds(self, tmp_path):
        ledger = _ledger(tmp_path)
        assert ledger.charge_scan("g1", "Test Group") is True
        ledger.close()

    def test_shadow_logs_would_charge(self, tmp_path, caplog):
        ledger = _ledger(tmp_path)
        with caplog.at_level(logging.INFO):
            ledger.charge_scan("g1", "Test Group")
        assert any("would_charge" in r.message for r in caplog.records)
        ledger.close()

    def test_shadow_tracks_cycle_total(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_scan("g1", "Group 1")
        ledger.charge_scan("g2", "Group 2")
        assert ledger.cycle_total == 1.0  # 2 × 0.5
        ledger.close()

    def test_shadow_lead_67_charges_5(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_lead("Dopasowano", "c1", "p1")
        assert ledger.cycle_total == 5.0
        ledger.close()

    def test_shadow_lead_68_charges_2(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_lead("Do weryfikacji", "c1", "p1")
        assert ledger.cycle_total == 2.0
        ledger.close()

    def test_shadow_lead_69_charges_0(self, tmp_path):
        ledger = _ledger(tmp_path)
        result = ledger.charge_lead("Brak dopasowania", "c1", "p1")
        assert result is True
        assert ledger.cycle_total == 0.0
        ledger.close()

    def test_shadow_never_calls_api(self, tmp_path):
        ledger = _ledger(tmp_path)
        # Patch the client to detect any calls
        ledger._client = MagicMock()
        ledger.charge_scan("g1", "Group")
        ledger.charge_lead("Dopasowano", "c1", "p1")
        ledger._client.charge.assert_not_called()
        ledger._client.check_balance.assert_not_called()
        ledger.close()

    def test_shadow_writes_billing_log(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_scan("g1", "Test Group")
        cur = ledger._conn.execute("SELECT * FROM billing_log")
        rows = cur.fetchall()
        assert len(rows) == 1
        assert rows[0]["result"] == "would_charge"
        assert rows[0]["mode"] == "shadow"
        assert rows[0]["tokens"] == 0.5
        assert rows[0]["reference"] == "scan-g1-1"
        ledger.close()


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

class TestLedgerDedup:
    def test_same_reference_not_charged_twice(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_scan("g1", "Group")
        ledger.charge_scan("g1", "Group")  # same group+cycle = same reference
        assert ledger.cycle_total == 0.5  # only charged once
        ledger.close()

    def test_different_references_both_charged(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_scan("g1", "Group 1")
        ledger.charge_scan("g2", "Group 2")
        assert ledger.cycle_total == 1.0
        ledger.close()

    def test_dedup_persists_across_instances(self, tmp_path):
        db = str(tmp_path / "test.db")
        s = _settings(db_path=db)

        l1 = BillingLedger(s, cycle_id=1)
        l1.charge_scan("g1", "Group")
        l1.close()

        l2 = BillingLedger(s, cycle_id=1)
        l2.charge_scan("g1", "Group")  # same reference
        assert l2.cycle_total == 0.0  # not charged again
        l2.close()

    def test_different_cycle_new_reference(self, tmp_path):
        db = str(tmp_path / "test.db")
        s = _settings(db_path=db)

        l1 = BillingLedger(s, cycle_id=1)
        l1.charge_scan("g1", "Group")
        l1.close()

        l2 = BillingLedger(s, cycle_id=2)
        l2.charge_scan("g1", "Group")  # different cycle = different reference
        assert l2.cycle_total == 0.5
        l2.close()


# ---------------------------------------------------------------------------
# Circuit breaker
# ---------------------------------------------------------------------------

class TestLedgerCircuitBreaker:
    def test_halts_when_exceeds_max(self, tmp_path):
        ledger = _ledger(tmp_path, billing_max_tokens_per_cycle=3.0)
        # 6 scans × 0.5 = 3.0 → 7th should trigger halt
        for i in range(6):
            ledger.charge_scan(f"g{i}", f"Group {i}")
        assert ledger.halted is False
        assert ledger.cycle_total == 3.0

        # Next charge exceeds limit
        result = ledger.charge_scan("g6", "Group 6")
        assert ledger.halted is True
        assert result is True  # shadow: still proceeds
        ledger.close()

    def test_halt_blocks_in_live(self, tmp_path):
        ledger = _ledger(tmp_path, billing_mode="live", billing_max_tokens_per_cycle=1.0)
        # Mock the client to avoid real API calls
        mock_client = MagicMock()
        mock_client.charge.return_value = ChargeResult(
            success=True, charged=0.5, balance_after=99.5, error=""
        )
        ledger._client = mock_client

        ledger.charge_scan("g1", "Group 1")  # 0.5 ✓
        ledger.charge_scan("g2", "Group 2")  # 1.0 ✓
        result = ledger.charge_scan("g3", "Group 3")  # 1.5 > 1.0 → halt
        assert ledger.halted is True
        assert result is False  # live: blocks action
        ledger.close()


# ---------------------------------------------------------------------------
# Live mode
# ---------------------------------------------------------------------------

class TestLedgerLive:
    def test_live_calls_api(self, tmp_path):
        ledger = _ledger(tmp_path, billing_mode="live")
        mock_client = MagicMock()
        mock_client.charge.return_value = ChargeResult(
            success=True, charged=0.5, balance_after=99.5, error=""
        )
        ledger._client = mock_client

        result = ledger.charge_scan("g1", "Group 1")
        assert result is True
        mock_client.charge.assert_called_once()
        # Verify reference format
        ref = mock_client.charge.call_args.args[2]
        assert ref == "scan-g1-1"
        ledger.close()

    def test_live_insufficient_blocks(self, tmp_path):
        ledger = _ledger(tmp_path, billing_mode="live")
        mock_client = MagicMock()
        mock_client.charge.side_effect = InsufficientTokensError("no funds")
        ledger._client = mock_client

        result = ledger.charge_scan("g1", "Group 1")
        assert result is False  # action blocked
        ledger.close()

    def test_live_api_error_blocks(self, tmp_path, caplog):
        ledger = _ledger(tmp_path, billing_mode="live")
        mock_client = MagicMock()
        mock_client.charge.side_effect = RuntimeError("timeout")
        ledger._client = mock_client

        with caplog.at_level(logging.ERROR):
            result = ledger.charge_scan("g1", "Group 1")
        assert result is False  # fail-safe: block
        assert any("airminal API error" in r.message for r in caplog.records)
        ledger.close()

    def test_live_no_key_blocks(self, tmp_path):
        ledger = _ledger(tmp_path, billing_mode="live", airminal_billing_key="")
        result = ledger.charge_scan("g1", "Group 1")
        assert result is False
        ledger.close()

    def test_live_logs_charged(self, tmp_path):
        ledger = _ledger(tmp_path, billing_mode="live")
        mock_client = MagicMock()
        mock_client.charge.return_value = ChargeResult(
            success=True, charged=5.0, balance_after=95.0, error=""
        )
        ledger._client = mock_client

        ledger.charge_lead("Dopasowano", "c1", "p1")
        cur = ledger._conn.execute("SELECT * FROM billing_log WHERE result='charged'")
        rows = cur.fetchall()
        assert len(rows) == 1
        assert rows[0]["tokens"] == 5.0
        assert rows[0]["balance_after"] == 95.0
        assert rows[0]["reference"] == "lead-c1-p1"
        ledger.close()


# ---------------------------------------------------------------------------
# Invalidate empty scans
# ---------------------------------------------------------------------------

class TestInvalidateEmptyScans:
    def test_invalidates_scans_from_empty_cycles(self, tmp_path):
        db = str(tmp_path / "test.db")
        s = _settings(db_path=db)

        # Create cycle_log table and two cycles
        import sqlite3
        conn = sqlite3.connect(db)
        conn.execute("""
            CREATE TABLE IF NOT EXISTS cycle_log (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                started_at TEXT NOT NULL,
                posts_fetched INTEGER DEFAULT 0
            )
        """)
        conn.execute("INSERT INTO cycle_log (started_at, posts_fetched) VALUES ('2026-01-01', 10)")
        conn.execute("INSERT INTO cycle_log (started_at, posts_fetched) VALUES ('2026-01-02', 0)")
        conn.commit()
        conn.close()

        # Create billing entries for both cycles
        l1 = BillingLedger(s, cycle_id=1)
        l1.charge_scan("g1", "Group 1")
        l1.close()

        l2 = BillingLedger(s, cycle_id=2)
        l2.charge_scan("g2", "Group 2")
        result = l2.invalidate_empty_scans(db)
        l2.close()

        assert result["invalidated"] == 1
        assert result["tokens"] == 0.5

        # Verify the entry is marked
        conn = sqlite3.connect(db)
        row = conn.execute(
            "SELECT result FROM billing_log WHERE cycle_id = '2'"
        ).fetchone()
        assert row[0] == "invalid_empty_scan"

        # Cycle 1 entry should be untouched
        row = conn.execute(
            "SELECT result FROM billing_log WHERE cycle_id = '1'"
        ).fetchone()
        assert row[0] == "would_charge"
        conn.close()

    def test_no_empty_cycles_returns_zero(self, tmp_path):
        db = str(tmp_path / "test.db")

        import sqlite3
        conn = sqlite3.connect(db)
        conn.execute("""
            CREATE TABLE IF NOT EXISTS cycle_log (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                started_at TEXT NOT NULL,
                posts_fetched INTEGER DEFAULT 0
            )
        """)
        conn.execute("INSERT INTO cycle_log (started_at, posts_fetched) VALUES ('2026-01-01', 5)")
        conn.commit()
        conn.close()

        s = _settings(db_path=db)
        ledger = BillingLedger(s, cycle_id=1)
        result = ledger.invalidate_empty_scans(db)
        ledger.close()
        assert result["invalidated"] == 0


# ---------------------------------------------------------------------------
# Reference format
# ---------------------------------------------------------------------------

class TestReferenceFormat:
    def test_scan_reference(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_scan("group_abc", "Test")
        cur = ledger._conn.execute("SELECT reference FROM billing_log")
        assert cur.fetchone()[0] == "scan-group_abc-1"
        ledger.close()

    def test_lead_reference(self, tmp_path):
        ledger = _ledger(tmp_path)
        ledger.charge_lead("Dopasowano", "client_42", "post_xyz")
        cur = ledger._conn.execute("SELECT reference FROM billing_log")
        assert cur.fetchone()[0] == "lead-client_42-post_xyz"
        ledger.close()


# ---------------------------------------------------------------------------
# API key masking — key must never appear in logs
# ---------------------------------------------------------------------------

class TestKeyMasking:
    def test_key_not_in_shadow_logs(self, tmp_path, caplog):
        ledger = _ledger(tmp_path, airminal_billing_key="SECRET_KEY_123")
        with caplog.at_level(logging.DEBUG):
            ledger.charge_scan("g1", "Group")
        all_logs = " ".join(r.message for r in caplog.records)
        assert "SECRET_KEY_123" not in all_logs
        ledger.close()

    def test_key_not_in_billing_log_table(self, tmp_path):
        ledger = _ledger(tmp_path, airminal_billing_key="SECRET_KEY_123")
        ledger.charge_scan("g1", "Group")
        cur = ledger._conn.execute("SELECT * FROM billing_log")
        for row in cur.fetchall():
            for col in row.keys():
                assert "SECRET_KEY_123" not in str(row[col])
        ledger.close()
