from __future__ import annotations from datetime import datetime, timedelta, timezone import pytest from sqlalchemy import select 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, close_reason: str = "stop", 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=close_reason, 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) async def test_same_day_fail_then_qualify_stays_locked(session): """Fail at 10:00 NY and qualify at 15:35 NY same day must not unlock.""" session.add(User(id=1, username="u", password_hash="x", role="user", has_access=True)) ticker = Ticker(symbol="SAMEDAY") session.add(ticker) await session.flush() stopped_at = datetime(2026, 7, 15, 14, 0, tzinfo=timezone.utc) # 10:00 ET session.add(_stopped_trade(ticker.id, closed_at=stopped_at)) await session.commit() fail_at = datetime(2026, 7, 15, 14, 5, tzinfo=timezone.utc) # ~10:05 ET updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={ticker.id}, qualified_ticker_ids=set(), observed_at=fail_at, ) assert updated == {ticker.id} qualify_same_day = datetime(2026, 7, 15, 19, 35, tzinfo=timezone.utc) # 15:35 ET updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={ticker.id}, qualified_ticker_ids={ticker.id}, observed_at=qualify_same_day, ) assert updated == set() assert ticker.id in await get_reentry_gate_locks(session) trade = ( await session.execute(select(PaperTrade).where(PaperTrade.ticker_id == ticker.id)) ).scalar_one() assert trade.reentry_gate_failed_at is not None assert trade.reentry_gate_requalified_at is None qualify_next_day = datetime(2026, 7, 16, 19, 35, tzinfo=timezone.utc) # next day 15:35 ET updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={ticker.id}, qualified_ticker_ids={ticker.id}, observed_at=qualify_next_day, ) assert updated == {ticker.id} await session.refresh(trade) assert trade.reentry_gate_requalified_at is not None assert ticker.id not in await get_reentry_gate_locks(session) async def test_newer_non_stop_exit_supersedes_historical_stop(session): session.add(User(id=1, username="u", password_hash="x", role="user", has_access=True)) ticker = Ticker(symbol="LATEREXIT") session.add(ticker) await session.flush() stopped_at = datetime.now(timezone.utc) - timedelta(days=20) old_stop = _stopped_trade(ticker.id, closed_at=stopped_at) later_manual_exit = _stopped_trade( ticker.id, closed_at=stopped_at + timedelta(days=10), close_reason="manual", ) session.add_all([old_stop, later_manual_exit]) await session.commit() assert ticker.id not in await get_reentry_gate_locks(session) observed_at = datetime.now(timezone.utc) updated = await observe_reentry_gate_transitions( session, evaluated_ticker_ids={ticker.id}, qualified_ticker_ids=set(), observed_at=observed_at, ) assert updated == set() assert old_stop.reentry_gate_failed_at is None