from __future__ import annotations from datetime import datetime, timedelta, timezone import pytest from app.models.paper_trade import PaperTrade from app.models.ticker import Ticker from app.models.user import User from app.services.trade_policy import ( get_reentry_gate_locks, observe_reentry_gate_transitions, ) from tests.conftest import _test_session_factory # type: ignore @pytest.fixture async def session(): async with _test_session_factory() as db: yield db def _stopped_trade( ticker_id: int, *, closed_at: datetime, gate_failed_at: datetime | None = None, gate_requalified_at: datetime | None = None, ) -> PaperTrade: return PaperTrade( user_id=1, ticker_id=ticker_id, direction="long", entry_price=100.0, shares=10.0, stop_loss=95.0, target=115.0, status="closed", opened_at=closed_at - timedelta(days=5), close_price=95.0, closed_at=closed_at, close_reason="stop", reentry_gate_failed_at=gate_failed_at, reentry_gate_requalified_at=gate_requalified_at, ) async def test_observation_releases_only_evaluated_unqualified_tickers(session): session.add(User(id=1, username="u", password_hash="x", role="user", has_access=True)) tickers = [ Ticker(symbol=symbol) for symbol in ("FAILQ", "PASSQ", "ERRORQ", "LATEQ") ] session.add_all(tickers) await session.flush() stopped_at = datetime.now(timezone.utc) - timedelta(days=1) trades = [ _stopped_trade(ticker.id, closed_at=stopped_at) for ticker in tickers[:3] ] observed_at = datetime.now(timezone.utc) trades.append( _stopped_trade( tickers[3].id, closed_at=observed_at + timedelta(seconds=1), ) ) session.add_all(trades) await session.commit() updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={tickers[0].id, tickers[1].id, tickers[3].id}, qualified_ticker_ids={tickers[1].id}, observed_at=observed_at, ) assert updated == {tickers[0].id} locks = await get_reentry_gate_locks(session) assert set(locks) == {ticker.id for ticker in tickers} assert trades[0].reentry_gate_failed_at == observed_at assert trades[0].reentry_gate_requalified_at is None assert trades[1].reentry_gate_failed_at is None assert trades[2].reentry_gate_failed_at is None assert trades[3].reentry_gate_failed_at is None requalified_at = observed_at + timedelta(days=1) updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={tickers[0].id}, qualified_ticker_ids={tickers[0].id}, observed_at=requalified_at, ) assert updated == {tickers[0].id} assert trades[0].reentry_gate_requalified_at == requalified_at assert set(await get_reentry_gate_locks(session)) == { tickers[1].id, tickers[2].id, tickers[3].id, } async def test_latest_stop_starts_a_new_gate_reset_episode(session): session.add(User(id=1, username="u", password_hash="x", role="user", has_access=True)) ticker = Ticker(symbol="TWOSTOP") session.add(ticker) await session.flush() first_stop = datetime.now(timezone.utc) - timedelta(days=20) session.add_all( [ _stopped_trade( ticker.id, closed_at=first_stop, gate_failed_at=first_stop + timedelta(days=1), gate_requalified_at=first_stop + timedelta(days=2), ), _stopped_trade( ticker.id, closed_at=first_stop + timedelta(days=10), ), ] ) await session.commit() assert ticker.id in await get_reentry_gate_locks(session)