from __future__ import annotations

from dataclasses import dataclass
from decimal import Decimal
from sqlite3 import Connection

from jarvis_finance.crypto.holdings import calculate_crypto_holdings


@dataclass(frozen=True)
class CryptoBalanceBasis:
    quantities: dict[tuple[str, str], Decimal]
    snapshot_id: str | None
    balance_as_of: str | None
    confirmed_current: bool


def current_crypto_balance_basis(conn: Connection) -> CryptoBalanceBasis:
    """Return the latest complete observed snapshot, or the legacy confirmed anchor."""
    table = conn.execute(
        "SELECT 1 FROM sqlite_master WHERE type='table' AND name='crypto_balance_snapshots'"
    ).fetchone()
    if table:
        snapshot = conn.execute(
            """SELECT snapshot_id, observed_at FROM crypto_balance_snapshots
               WHERE status='complete' ORDER BY observed_at DESC, created_at DESC LIMIT 1"""
        ).fetchone()
        if snapshot:
            rows = conn.execute(
                """SELECT wallet_id, asset_id, quantity
                   FROM crypto_balance_snapshot_items WHERE snapshot_id=?""",
                (snapshot["snapshot_id"],),
            ).fetchall()
            quantities = {(row["wallet_id"], row["asset_id"]): Decimal(str(row["quantity"])) for row in rows}
            balance_as_of = str(snapshot["observed_at"])
            later = conn.execute(
                """SELECT asset_id,quantity,fee_quantity,from_wallet_id,to_wallet_id,transaction_datetime
                   FROM crypto_transactions
                   WHERE confirmation_status='confirmed' AND transaction_datetime>?
                   ORDER BY transaction_datetime,created_at,crypto_transaction_id""",
                (balance_as_of,),
            ).fetchall()
            for row in later:
                asset_id = str(row["asset_id"])
                quantity = Decimal(str(row["quantity"] or "0"))
                fee = Decimal(str(row["fee_quantity"] or "0"))
                if row["from_wallet_id"]:
                    key = (str(row["from_wallet_id"]), asset_id)
                    quantities[key] = quantities.get(key, Decimal("0")) - quantity - fee
                if row["to_wallet_id"]:
                    key = (str(row["to_wallet_id"]), asset_id)
                    quantities[key] = quantities.get(key, Decimal("0")) + quantity
                balance_as_of = max(balance_as_of, str(row["transaction_datetime"]))
            return CryptoBalanceBasis(
                quantities=quantities,
                snapshot_id=str(snapshot["snapshot_id"]),
                balance_as_of=balance_as_of,
                confirmed_current=True,
            )
    holdings = calculate_crypto_holdings(conn)
    anchor = conn.execute(
        "SELECT MAX(COALESCE(last_verified_at, legacy_snapshot_date)) FROM crypto_holdings"
    ).fetchone()[0]
    return CryptoBalanceBasis(
        quantities={key: item.quantity for key, item in holdings.wallet_holdings.items()},
        snapshot_id=None,
        balance_as_of=str(anchor) if anchor else None,
        confirmed_current=False,
    )
