import pytest from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine import app.bets.confirmation # noqa: F401 (registers the "bet" handler) import app.rounds.confirmation # noqa: F401 (registers the "payout" handler) from sqlalchemy import select from app.config import settings from app.db.base import Base from app.db.models import PendingTransaction, Round, RoundParticipant, User from app.electrum.scripthash import address_to_scripthash from app.tx.confirmation import poll_once from app.wallet.hd import derive_user_address class FakeClient: """B-41: poll_once now asks blockchain.scripthash.get_history rather than a verbose blockchain.transaction.get, so this hands back a flat history — height > 0 means confirmed at that height, 0 (or absent) means still in the mempool. The scripthash argument is ignored: every candidate's derived address is looked up against the same known universe of txids, which is fine since matching happens on tx_hash, not on which address asked.""" def __init__(self, heights_by_txid: dict[str, int]): self._heights = heights_by_txid async def get_history(self, scripthash: str) -> list[dict]: return [{"tx_hash": txid, "height": height} for txid, height in self._heights.items()] @pytest.fixture async def session_factory(tmp_path, monkeypatch): # own_address_for (B-41) derives each row's address via the HD wallet, so # poll_once now needs a real master key — same bootstrap test_broadcast.py # and test_reconcile.py use. monkeypatch.setattr(settings, "master_key_path", str(tmp_path / "master.xprv.enc")) monkeypatch.setattr( settings, "xprv_encryption_key", __import__("cryptography.fernet", fromlist=["Fernet"]).Fernet.generate_key().decode(), ) from app.wallet import hd hd._account_key = None hd.generate_master_key() engine = create_async_engine("sqlite+aiosqlite:///:memory:") async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) yield async_sessionmaker(engine, expire_on_commit=False) await engine.dispose() hd._account_key = None async def _make_user(session, derivation_index: int) -> User: user = User( username=f"user{derivation_index}", password_hash="x", derivation_index=derivation_index, address=derive_user_address(derivation_index), ) session.add(user) await session.flush() return user async def test_bet_confirmation_marks_participant_confirmed(session_factory): async with session_factory() as session: user = await _make_user(session, 0) session.add(Round(id=1, status="open")) session.add( RoundParticipant( round_id=1, user_id=user.id, bet_amount_sats=1_000, bet_txid="tx1", status="broadcast" ) ) session.add( PendingTransaction( kind="bet", round_id=1, user_id=user.id, current_txid="tx1", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) await session.commit() client = FakeClient({"tx1": 1}) confirmed = await poll_once(session_factory, client) assert confirmed == 1 async with session_factory() as session: participant = (await session.scalars(select(RoundParticipant))).one() assert participant.status == "confirmed" assert participant.confirmed_at is not None pending = (await session.scalars(select(PendingTransaction))).one() assert pending.status == "confirmed" async def test_unconfirmed_tx_is_left_pending(session_factory): async with session_factory() as session: user = await _make_user(session, 0) session.add(Round(id=2, status="open")) session.add( RoundParticipant(round_id=2, user_id=user.id, bet_amount_sats=1_000, bet_txid="tx2", status="broadcast") ) session.add( PendingTransaction( kind="bet", round_id=2, user_id=user.id, current_txid="tx2", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) await session.commit() client = FakeClient({"tx2": 0}) confirmed = await poll_once(session_factory, client) assert confirmed == 0 async with session_factory() as session: participant = (await session.scalars(select(RoundParticipant))).one() assert participant.status == "broadcast" async def test_payout_confirmation_closes_round(session_factory): async with session_factory() as session: session.add(Round(id=3, status="paying_out", payout_txid="tx3")) session.add( PendingTransaction( kind="payout", round_id=3, current_txid="tx3", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending" ) ) await session.commit() client = FakeClient({"tx3": 2}) confirmed = await poll_once(session_factory, client) assert confirmed == 1 async with session_factory() as session: round_ = await session.get(Round, 3) assert round_.status == "closed" class ExplodingClient: """Answers for one address's history and raises for the other's — the get_history equivalent of a server that no longer knows a particular tx (dropped from the mempool, replaced by a bump).""" def __init__(self, heights_by_txid: dict[str, int], exploding_scripthash: str): self._heights = heights_by_txid self._exploding = exploding_scripthash async def get_history(self, scripthash: str) -> list[dict]: if scripthash == self._exploding: raise RuntimeError("server error") return [{"tx_hash": txid, "height": height} for txid, height in self._heights.items()] async def test_one_unresolvable_candidate_does_not_block_the_others(session_factory): """B-03: the lookup used to be unguarded, so a single failing candidate aborted the whole pass — nothing confirmed again until an operator intervened, which in turn meant no round could ever close. B-41 changed the failure unit from "one txid" to "one address's history", but the isolation guarantee is the same.""" async with session_factory() as session: good_user = await _make_user(session, 0) gone_user = await _make_user(session, 1) session.add(Round(id=10, status="open")) session.add( RoundParticipant( round_id=10, user_id=good_user.id, bet_amount_sats=1_000, bet_txid="good", status="broadcast" ) ) session.add( PendingTransaction( kind="bet", round_id=10, user_id=gone_user.id, current_txid="gone", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) session.add( PendingTransaction( kind="bet", round_id=10, user_id=good_user.id, current_txid="good", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) await session.commit() exploding_scripthash = address_to_scripthash(derive_user_address(1)) confirmed = await poll_once( session_factory, ExplodingClient({"good": 1}, exploding_scripthash=exploding_scripthash) ) assert confirmed == 1 # the healthy one still got processed async with session_factory() as session: participant = (await session.scalars(select(RoundParticipant).where(RoundParticipant.round_id == 10))).one() assert participant.status == "confirmed" rows = {p.current_txid: p.status for p in (await session.scalars(select(PendingTransaction))).all()} assert rows["good"] == "confirmed" assert rows["gone"] == "pending" # left for the reconciler to judge, not abandoned here async def test_bet_confirms_after_an_rbf_bump_changed_the_txid(session_factory): """B-02: the handler used to match on bet_txid, so a bumped bet confirmed under a txid no participant carried — the participant stayed "broadcast" forever and the round could never close. It now resolves by (round_id, user_id).""" async with session_factory() as session: user = await _make_user(session, 0) session.add(Round(id=11, status="open")) session.add( RoundParticipant( round_id=11, user_id=user.id, bet_amount_sats=1_000, bet_txid="old-txid", status="broadcast" ) ) session.add( PendingTransaction( kind="bet", round_id=11, user_id=user.id, current_txid="bumped-txid", fee_rate_sat_vb=2, raw_tx_hex="00", status="pending", replaced_by_txid="old-txid", ) ) await session.commit() assert await poll_once(session_factory, FakeClient({"bumped-txid": 1})) == 1 async with session_factory() as session: participant = (await session.scalars(select(RoundParticipant).where(RoundParticipant.round_id == 11))).one() assert participant.status == "confirmed" async def test_payout_confirms_after_an_rbf_bump_changed_the_txid(session_factory): async with session_factory() as session: session.add(Round(id=12, status="paying_out", payout_txid="old-payout")) session.add( PendingTransaction( kind="payout", round_id=12, current_txid="bumped-payout", fee_rate_sat_vb=2, raw_tx_hex="00", status="pending", ) ) await session.commit() assert await poll_once(session_factory, FakeClient({"bumped-payout": 1})) == 1 async with session_factory() as session: assert (await session.get(Round, 12)).status == "closed" async def test_poll_once_caches_history_per_scripthash(session_factory): """Two pending bets from the same user share one address — fetching its history twice in one pass would be wasteful.""" async with session_factory() as session: user = await _make_user(session, 0) session.add(Round(id=20, status="open")) session.add( PendingTransaction( kind="bet", round_id=20, user_id=user.id, current_txid="tx-a", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) session.add( PendingTransaction( kind="withdrawal", user_id=user.id, current_txid="tx-b", fee_rate_sat_vb=1, raw_tx_hex="00", status="pending", ) ) await session.commit() call_count = {"n": 0} class CountingClient: async def get_history(self, scripthash: str) -> list[dict]: call_count["n"] += 1 return [{"tx_hash": "tx-a", "height": 0}, {"tx_hash": "tx-b", "height": 0}] await poll_once(session_factory, CountingClient()) assert call_count["n"] == 1