"""ASEx AccountProjectionV2 wrapper tests — race-proof, stress, seam.""" from __future__ import annotations import gc import math import threading import time from concurrent.futures import ThreadPoolExecutor, as_completed from types import SimpleNamespace import pytest from .account import AccountProjectionV2, AccountSnapshotV2, EPosition, TradeStage, TradeSide from .asex_account import ASEXAccountV2, _AccountBackend from asex.guarded import ValidationError def _empty_slots(n=1): return [SimpleNamespace( slot_id=i, trade_id=f"t{i}", asset="BTCUSDT", side=TradeSide.LONG, entry_price=0.0, size=0.0, initial_size=0.0, leverage=1.0, realized_pnl=0.0, closed=False, fsm_state=TradeStage.IDLE, exit_leg_ratios=(1.0,), active_leg_index=0, active_exit_order=None, active_entry_order=None, close_reason="", entry_time=None, last_event_time=None, seen_event_ids=(), metadata={}, unrealized_pnl=0.0, to_dict=lambda: {}, ) for i in range(n)] def _clean(): gc.collect(); gc.collect() def _force_race(obj, attr, n=10): b = threading.Barrier(n) def _w(): b.wait(); v = getattr(obj, attr); b.wait(); setattr(obj, attr, v + 1) ts = [threading.Thread(target=_w) for _ in range(n)] for t in ts: t.start() for t in ts: t.join(timeout=10) return getattr(obj, attr) class TestRaceProof: def test_k_realized_races(self): assert _force_race(AccountProjectionV2(0.0), "_k_realized", 50) == 1 def test_k_fees_races(self): assert _force_race(AccountProjectionV2(0.0), "_k_fees", 50) == 1 def test_k_funding_races(self): assert _force_race(AccountProjectionV2(0.0), "_k_funding", 50) == 1 def test_event_seq_races(self): assert _force_race(AccountProjectionV2(0.0), "_event_seq", 50) == 1 class TestASExBasic: def test_seed_only(self): _clean(); p = ASEXAccountV2(10000.0) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.capital == pytest.approx(10000.0); p.close(); _clean() def test_realized_adds(self): _clean(); p = ASEXAccountV2(10000.0) p.apply_fill(fill_price=100, fill_qty=1, fee=0, realized_pnl=500) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.capital == pytest.approx(10500.0); p.close(); _clean() def test_fee_subtracts(self): _clean(); p = ASEXAccountV2(10000.0) p.apply_fill(fill_price=100, fill_qty=1, fee=3.5, realized_pnl=0) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.capital == pytest.approx(9996.5); p.close(); _clean() def test_funding_subtracts(self): _clean(); p = ASEXAccountV2(10000.0) p.apply_funding(7.25) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.capital == pytest.approx(9992.75); p.close(); _clean() def test_combined(self): _clean(); p = ASEXAccountV2(10000.0) p.apply_fill(fill_price=50, fill_qty=2, fee=2, realized_pnl=100) p.apply_funding(5) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.capital == pytest.approx(10093.0) p.close(); _clean() class TestASExConcurrency: @pytest.mark.parametrize("n,ops", [(10, 100), (20, 100), (50, 50)]) def test_no_lost_updates(self, n, ops): _clean(); p = ASEXAccountV2(0.0) def _w(tid): for i in range(ops): p.apply_fill(fill_price=100, fill_qty=1, fee=0, realized_pnl=float(tid * ops + i)) with ThreadPoolExecutor(max_workers=n) as ex: for f in as_completed([ex.submit(_w, i) for i in range(n)]): f.result(timeout=60) expected = sum(tid * ops + i for tid in range(n) for i in range(ops)) assert p._backend._proj._k_realized == pytest.approx(float(expected)) p.close(); _clean() def test_event_seq_monotonic(self): _clean(); p = ASEXAccountV2(0.0); seqs = [] def _w(tid): for i in range(50): s = p.build_snapshot(source_event_id=f"e{tid}_{i}", slots=_empty_slots(), ts=float(i)) seqs.append(s.event_seq) with ThreadPoolExecutor(max_workers=10) as ex: for f in as_completed([ex.submit(_w, i) for i in range(10)]): f.result(timeout=60) assert sorted(seqs) == list(range(1, 501)) p.close(); _clean() def test_mixed_ops(self): _clean(); p = ASEXAccountV2(10000.0) def _f(tid): for i in range(100): p.apply_fill(fill_price=float(i), fill_qty=1, fee=0, realized_pnl=1) def _fu(tid): for i in range(50): p.apply_funding(0.5) with ThreadPoolExecutor(max_workers=4) as ex: for f in as_completed([ex.submit(_f, i) for i in [0,1]] + [ex.submit(_fu, i) for i in [0,1]]): f.result(timeout=60) assert p._backend._proj._k_realized == pytest.approx(200.0) # 2 workers x 50 calls x 0.5 = 50.0; exact total proves zero lost updates assert p._backend._proj._k_funding == pytest.approx(50.0) p.close(); _clean() class TestASExSeam: def test_backend_rejects_no_op(self): assert not _AccountBackend(0.0)._validate({}) def test_backend_rejects_unknown(self): b = _AccountBackend(0.0) with pytest.raises(Exception): b.mutate({"op": "nope"}) def test_backend_applied(self): b = _AccountBackend(0.0) b.mutate({"op": "apply_fill", "args": {"fill_price": 100, "fill_qty": 1, "fee": 0, "realized_pnl": 10}}) assert b.applied == 1 def test_worker_alive(self): p = ASEXAccountV2(0.0); assert p._worker._worker.is_alive(); p.close() def test_double_close(self): p = ASEXAccountV2(0.0); p.close(); p.close() def test_context_manager(self): with ASEXAccountV2(10000.0) as p: p.apply_fill(fill_price=100, fill_qty=1, fee=0, realized_pnl=50) s = p.build_snapshot(source_event_id="t", slots=_empty_slots(), ts=1e6) assert s.k.realized_pnl == pytest.approx(50.0) def test_no_thread_leak(self): _clean(); b = threading.active_count() for _ in range(50): p = ASEXAccountV2(0.0); p.apply_fill(fill_price=100, fill_qty=1, fee=0, realized_pnl=1); p.close() _clean(); assert threading.active_count() - b <= 2 class TestLockProof: def test_k_realized_lock_proof(self): ap, l, n = AccountProjectionV2(0.0), threading.Lock(), 50 b = threading.Barrier(n) def _w(): b.wait() with l: ap._k_realized += 1 ts = [threading.Thread(target=_w) for _ in range(n)] for t in ts: t.start() for t in ts: t.join(timeout=10) assert ap._k_realized == n def test_event_seq_lock_proof(self): ap, l, n = AccountProjectionV2(0.0), threading.Lock(), 50 b = threading.Barrier(n) def _w(): b.wait() with l: ap._event_seq += 1 ts = [threading.Thread(target=_w) for _ in range(n)] for t in ts: t.start() for t in ts: t.join(timeout=10) assert ap._event_seq == n