from __future__ import annotations

import urllib.error
from datetime import datetime, timedelta, timezone
from decimal import Decimal

from jarvis_finance.cli.main import main
from jarvis_finance.crypto.assets import create_crypto_asset
from jarvis_finance.crypto.holdings import create_initial_holding_snapshot
from jarvis_finance.crypto.wallets import create_wallet
from jarvis_finance.dashboard import data
from jarvis_finance.market.providers import CoinGeckoClient, PriceQuote, refresh_crypto_prices
from jarvis_finance.reports.crypto_inventory import build_crypto_report_context
from jarvis_finance.storage.database import connect_memory
from jarvis_finance.storage.migrations import apply_migrations


class MockProvider:
    def __init__(self) -> None:
        self.calls: list[tuple[str, str]] = []
        self.batch_calls: list[tuple[tuple[str, ...], str]] = []

    def get_crypto_price(self, coingecko_id: str, currency: str = "CHF") -> PriceQuote:
        self.calls.append((coingecko_id, currency))
        return PriceQuote(coingecko_id=coingecko_id, currency=currency, price=Decimal("123.456789"), provider="MockGecko", provider_timestamp="2026-01-31T00:00:00+00:00")

    def get_crypto_prices(self, coingecko_ids, currency: str = "CHF"):
        self.batch_calls.append((tuple(coingecko_ids), currency))
        return {coingecko_id: self.get_crypto_price(coingecko_id, currency) for coingecko_id in coingecko_ids}


def setup_conn():
    conn = connect_memory()
    apply_migrations(conn)
    wallet_id = create_wallet(conn, wallet_name="Synthetic Wallet", wallet_type="Hardware Wallet", platform_provider="Synthetic", last_verified_at="2026-01-31T00:00:00+00:00")
    btc_id = create_crypto_asset(conn, coin_name="Bitcoin", symbol="BTC", coingecko_id="bitcoin")
    create_initial_holding_snapshot(conn, asset_id=btc_id, wallet_id=wallet_id, quantity=Decimal("0.01"), verification_status="verified", note="synthetic fixture", last_verified_at="2026-01-31T00:00:00+00:00")
    missing_id = create_crypto_asset(conn, coin_name="Missing Coin", symbol="MISS", coingecko_id=None)
    create_initial_holding_snapshot(conn, asset_id=missing_id, wallet_id=wallet_id, quantity=Decimal("1"), verification_status="verified", note="synthetic fixture", last_verified_at="2026-01-31T00:00:00+00:00")
    return conn, btc_id, missing_id


def test_refresh_crypto_prices_loads_active_assets_and_stores_chf_price() -> None:
    conn, btc_id, _ = setup_conn()
    provider = MockProvider()

    result = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=0)

    assert result.success_count == 1
    assert result.skipped_count == 1
    assert result.warning_count == 1
    assert result.error_count == 0
    assert provider.batch_calls == [(("bitcoin",), "CHF")]
    assert provider.calls == [("bitcoin", "CHF")]
    row = conn.execute("SELECT * FROM crypto_prices WHERE asset_id=?", (btc_id,)).fetchone()
    assert row is not None
    assert row["price"] == "123.456789"
    assert row["price_currency"] == "CHF"
    assert row["provider"] == "MockGecko"


def test_missing_coingecko_id_is_skipped_and_warned_without_duplicate_alerts() -> None:
    conn, _, missing_id = setup_conn()
    provider = MockProvider()

    first = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=0)
    second = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=0)

    assert first.skipped_count == 1
    assert second.skipped_count == 1
    alerts = conn.execute("SELECT * FROM alerts WHERE entity_id=? AND rule_id='missing_coingecko_id'", (missing_id,)).fetchall()
    assert len(alerts) == 1
    assert alerts[0]["occurrence_count"] >= 2


def test_429_backoff_is_used_by_provider() -> None:
    calls = {"n": 0}
    sleeps: list[float] = []

    def opener(url, timeout=20):
        calls["n"] += 1
        if calls["n"] == 1:
            raise urllib.error.HTTPError(url, 429, "rate limited", None, None)
        class Resp:
            def __enter__(self): return self
            def __exit__(self, *args): return False
            def read(self): return b'{"bitcoin":{"chf":"12.34","last_updated_at":1700000000}}'
        return Resp()

    client = CoinGeckoClient(max_retries=2, initial_backoff_seconds=0.25, max_backoff_seconds=1, opener=opener, sleeper=sleeps.append)

    quote = client.get_crypto_price("bitcoin", "CHF")

    assert quote.price == Decimal("12.34")
    assert sleeps == [0.25]
    assert calls["n"] == 2


def test_fresh_cache_prevents_refetch_and_stale_price_is_marked() -> None:
    conn, btc_id, _ = setup_conn()
    provider = MockProvider()
    first = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=3600)
    second = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=3600)
    assert first.success_count == 1
    assert second.cached_count == 1
    assert provider.calls == [("bitcoin", "CHF")]

    old = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
    conn.execute("UPDATE crypto_prices SET fetched_at=?, provider_timestamp=? WHERE asset_id=?", (old, old, btc_id))
    conn.commit()

    overview = data.get_crypto_overview(conn, price_max_age_seconds=3600)
    btc = next(row for row in overview if row["symbol"] == "BTC")
    assert btc["price_quality"] == "stale"
    context = build_crypto_report_context(conn, price_max_age_seconds=3600)
    assert any(w["code"] == "stale_price" and w["entity_id"] == btc_id for w in context["data_quality_warnings"])


def test_dashboard_reads_local_prices_without_provider_api_call() -> None:
    conn, btc_id, _ = setup_conn()
    refresh_crypto_prices(conn, provider=MockProvider(), currency="CHF", max_age_seconds=0)

    overview = data.get_crypto_overview(conn)
    by_wallet = data.get_crypto_by_wallet(conn)
    summary = data.get_command_center_summary(conn)

    btc = next(row for row in overview if row["symbol"] == "BTC")
    assert btc["latest_price_chf"] == "123.46"
    assert btc["market_value_chf"]
    assert any(row["wallet_value_chf"] for row in by_wallet if row["symbol"] == "BTC")
    assert summary["crypto_total_chf"] != "0.00"
    assert summary["crypto_price_data_as_of"]


def test_crypto_report_uses_local_prices_and_improves_missing_price_status() -> None:
    conn, btc_id, _ = setup_conn()
    before = build_crypto_report_context(conn)
    assert any(w["code"] == "missing_price" and w["entity_id"] == btc_id for w in before["data_quality_warnings"])

    refresh_crypto_prices(conn, provider=MockProvider(), currency="CHF", max_age_seconds=0)
    after = build_crypto_report_context(conn)

    btc = next(row for row in after["coins"] if row["symbol"] == "BTC")
    assert btc["price_chf"] == "123.456789"
    assert btc["value_chf"]
    assert not any(w["code"] == "missing_price" and w["entity_id"] == btc_id for w in after["data_quality_warnings"])
    assert after["metadata"]["live_api_calls"] is False


def test_cli_update_crypto_prices_command_uses_provider_factory(monkeypatch, tmp_path, capsys) -> None:
    runtime = tmp_path / "runtime"
    monkeypatch.setenv("JARVIS_FINANCE_RUNTIME_DIR", str(runtime))
    conn = connect_memory()
    # Build a runtime file DB via CLI settings by seeding after opening path.
    from jarvis_finance.storage.database import connect
    db = runtime / "data" / "finance.sqlite3"
    db.parent.mkdir(parents=True)
    file_conn = connect(db)
    apply_migrations(file_conn)
    wallet_id = create_wallet(file_conn, wallet_name="CLI Wallet", wallet_type="Hardware Wallet", platform_provider="Synthetic", last_verified_at="2026-01-31T00:00:00+00:00")
    asset_id = create_crypto_asset(file_conn, coin_name="Bitcoin", symbol="BTC", coingecko_id="bitcoin")
    create_initial_holding_snapshot(file_conn, asset_id=asset_id, wallet_id=wallet_id, quantity=Decimal("0.01"), verification_status="verified", note="synthetic fixture", last_verified_at="2026-01-31T00:00:00+00:00")
    file_conn.close()

    monkeypatch.setattr("jarvis_finance.cli.main.CoinGeckoClient", lambda **kwargs: MockProvider())

    rc = main(["update-crypto-prices", "--max-age-seconds", "0"])

    out = capsys.readouterr().out
    assert rc == 0
    assert "CRYPTO_PRICE_UPDATE" in out
    assert "updated=1" in out
    assert "assets_still_missing_local_price=0" in out
    conn.close()


def _add_asset_with_holding(conn, wallet_id: str, *, coin_name: str, symbol: str, coingecko_id: str):
    asset_id = create_crypto_asset(conn, coin_name=coin_name, symbol=symbol, coingecko_id=coingecko_id)
    create_initial_holding_snapshot(conn, asset_id=asset_id, wallet_id=wallet_id, quantity=Decimal("2"), verification_status="verified", note="synthetic fixture", last_verified_at="2026-01-31T00:00:00+00:00")
    return asset_id


def test_batch_request_updates_multiple_assets_in_one_provider_call() -> None:
    conn, btc_id, _ = setup_conn()
    wallet_id = conn.execute("SELECT wallet_id FROM crypto_wallets LIMIT 1").fetchone()["wallet_id"]
    eth_id = _add_asset_with_holding(conn, wallet_id, coin_name="Ethereum", symbol="ETH", coingecko_id="ethereum")
    provider = MockProvider()

    result = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=0)

    assert result.updated_count == 2
    assert provider.batch_calls == [(("bitcoin", "ethereum"), "CHF")]
    assert conn.execute("SELECT COUNT(*) AS n FROM crypto_prices WHERE asset_id IN (?, ?)", (btc_id, eth_id)).fetchone()["n"] == 2


def test_only_missing_refreshes_assets_without_any_fresh_price() -> None:
    conn, btc_id, _ = setup_conn()
    wallet_id = conn.execute("SELECT wallet_id FROM crypto_wallets LIMIT 1").fetchone()["wallet_id"]
    _add_asset_with_holding(conn, wallet_id, coin_name="Ethereum", symbol="ETH", coingecko_id="ethereum")
    refresh_crypto_prices(conn, provider=MockProvider(), currency="CHF", max_age_seconds=3600, only_symbol="BTC")
    provider = MockProvider()

    result = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=3600, only_missing=True)

    assert result.updated_count == 1
    assert provider.batch_calls == [(("ethereum",), "CHF")]
    assert conn.execute("SELECT COUNT(*) AS n FROM crypto_prices WHERE asset_id=?", (btc_id,)).fetchone()["n"] == 1


def test_only_stale_refreshes_only_stale_local_prices() -> None:
    conn, btc_id, _ = setup_conn()
    wallet_id = conn.execute("SELECT wallet_id FROM crypto_wallets LIMIT 1").fetchone()["wallet_id"]
    _add_asset_with_holding(conn, wallet_id, coin_name="Ethereum", symbol="ETH", coingecko_id="ethereum")
    refresh_crypto_prices(conn, provider=MockProvider(), currency="CHF", max_age_seconds=3600)
    old = (datetime.now(timezone.utc) - timedelta(days=10)).isoformat()
    conn.execute("UPDATE crypto_prices SET fetched_at=?, provider_timestamp=? WHERE asset_id=?", (old, old, btc_id))
    conn.commit()
    provider = MockProvider()

    result = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=3600, only_stale=True)

    assert result.updated_count == 1
    assert provider.batch_calls == [(("bitcoin",), "CHF")]


def test_only_symbol_and_limit_constrain_refresh_scope() -> None:
    conn, _, _ = setup_conn()
    wallet_id = conn.execute("SELECT wallet_id FROM crypto_wallets LIMIT 1").fetchone()["wallet_id"]
    _add_asset_with_holding(conn, wallet_id, coin_name="Ethereum", symbol="ETH", coingecko_id="ethereum")
    _add_asset_with_holding(conn, wallet_id, coin_name="Solana", symbol="SOL", coingecko_id="solana")
    provider = MockProvider()

    symbol_result = refresh_crypto_prices(conn, provider=provider, currency="CHF", max_age_seconds=0, only_symbol="ETH")
    assert symbol_result.updated_count == 1
    assert provider.batch_calls == [(("ethereum",), "CHF")]

    provider2 = MockProvider()
    limit_result = refresh_crypto_prices(conn, provider=provider2, currency="CHF", max_age_seconds=0, limit=1)
    assert limit_result.updated_count == 1
    assert len(provider2.batch_calls[0][0]) == 1


def test_dry_run_does_not_write_prices_or_alerts() -> None:
    conn, btc_id, missing_id = setup_conn()
    alerts_before = conn.execute("SELECT COUNT(*) AS n FROM alerts WHERE entity_id IN (?, ?)", (btc_id, missing_id)).fetchone()["n"]

    result = refresh_crypto_prices(conn, provider=MockProvider(), currency="CHF", max_age_seconds=0, dry_run=True)

    assert result.dry_run is True
    assert result.updated_count == 1
    assert result.missing_local_price_count == 2
    assert conn.execute("SELECT COUNT(*) AS n FROM crypto_prices").fetchone()["n"] == 0
    assert conn.execute("SELECT COUNT(*) AS n FROM alerts WHERE entity_id IN (?, ?)", (btc_id, missing_id)).fetchone()["n"] == alerts_before


def test_rate_limit_after_retries_marks_batch_stale_and_backoff_is_capped() -> None:
    sleeps: list[float] = []
    calls = {"n": 0}

    def opener(url, timeout=20):
        calls["n"] += 1
        raise urllib.error.HTTPError(url, 429, "rate limited", None, None)

    client = CoinGeckoClient(max_retries=2, initial_backoff_seconds=0.5, max_backoff_seconds=0.75, opener=opener, sleeper=sleeps.append)

    quotes = client.get_crypto_prices(["bitcoin", "ethereum"], "CHF")

    assert calls["n"] == 3
    assert sleeps == [0.5, 0.75]
    assert {quote.quality_status for quote in quotes.values()} == {"stale"}


def test_missing_local_price_alert_is_deduped_for_failed_refreshes() -> None:
    class MissingProvider:
        def get_crypto_prices(self, coingecko_ids, currency="CHF"):
            return {cid: PriceQuote(cid, currency, None, quality_status="missing", error_message="not listed") for cid in coingecko_ids}

    conn, btc_id, _ = setup_conn()
    refresh_crypto_prices(conn, provider=MissingProvider(), currency="CHF", max_age_seconds=0, only_symbol="BTC")
    refresh_crypto_prices(conn, provider=MissingProvider(), currency="CHF", max_age_seconds=0, only_symbol="BTC")

    alerts = conn.execute("SELECT * FROM alerts WHERE entity_id=? AND rule_id='crypto_price_missing_local'", (btc_id,)).fetchall()
    assert len(alerts) == 1
    assert alerts[0]["occurrence_count"] >= 2
