From 9a0d2ab2202f951d358ba9d6a6b7426625de6756 Mon Sep 17 00:00:00 2001 From: Le Date: Fri, 3 Jul 2026 22:29:41 +0700 Subject: [PATCH] fix: sua loi eviction dung nham gia cross-symbol trong trade_executor.py Phat hien (p) khi viet test cho trade_executor: khi kiem tra hybrid eviction, PnL cua TAT CA cac trade dang mo (o nhieu symbol khac nhau) bi tinh bang current_price cua tin hieu dang xu ly, thay vi gia thuc cua tung symbol. Fix bang cach lookup gia moi nhat theo tung symbol/exchange/timeframe (batched query, cung pattern da dung dung trong close_stale_trades), ap dung cho ca xep hang loser LAN gia dong lenh cuoi cung. Co-Authored-By: Claude Sonnet 5 --- backend/app/services/trade_executor.py | 43 +++++++++++- backend/tests/conftest.py | 94 +++++++++++++++++++++----- backend/tests/test_trade_executor.py | 68 ++++++++++++++++++- 3 files changed, 183 insertions(+), 22 deletions(-) diff --git a/backend/app/services/trade_executor.py b/backend/app/services/trade_executor.py index 97add5e..416b60a 100644 --- a/backend/app/services/trade_executor.py +++ b/backend/app/services/trade_executor.py @@ -15,8 +15,11 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.core.exceptions import AppException from app.database import async_session_factory +from app.models.candle import Candle +from app.models.exchange import Exchange from app.models.real_trade import RealTrade from app.models.signal import HypotheticalTrade, Signal +from app.models.symbol import Symbol from app.models.user import User from app.services.audit_service import log_action @@ -195,9 +198,42 @@ async def execute_signal_trade( if open_count >= MAX_OPEN_TRADES: to_evict = open_count - MAX_OPEN_TRADES + 1 + + # Each open trade may be on a different symbol than the signal + # currently being processed — using `current_price` (that + # signal's price) for all of them would rank PnL against the + # wrong market price. Look up each trade's own latest candle + # price instead (batched by unique symbol/exchange/timeframe). + trade_keys = list({(t.symbol, t.exchange, t.timeframe) for t in all_open_trades}) + price_map: dict[tuple[str, str, str], Decimal] = {} + for sym, ex, tf in trade_keys: + price_result = await db.execute( + select(Candle.close) + .select_from(Symbol) + .join(Candle, Candle.symbol_id == Symbol.id) + .join(Exchange, Exchange.id == Symbol.exchange_id) + .where(and_( + Exchange.name == ex, + Symbol.symbol == sym, + Candle.timeframe == tf, + )) + .order_by(desc(Candle.timestamp)) + .limit(1) + ) + row = price_result.first() + if row: + price_map[(sym, ex, tf)] = row[0] + + def _price_for(t: HypotheticalTrade) -> Decimal: + if t.symbol == symbol and t.exchange == exchange_name: + return current_price + # Fall back to entry_price (PnL=0, neutral) if no candle + # data is available for this trade's own symbol. + return price_map.get((t.symbol, t.exchange, t.timeframe), t.entry_price) + open_with_pnl = [] for t in all_open_trades: - pnl_val, _pct = _calculate_pnl(t.entry_price, current_price, t.direction, t.quantity) + pnl_val, _pct = _calculate_pnl(t.entry_price, _price_for(t), t.direction, t.quantity) open_with_pnl.append((t, pnl_val)) losers = [(t, pnl) for t, pnl in open_with_pnl if pnl < 0] @@ -208,10 +244,11 @@ async def execute_signal_trade( eviction_candidates = all_open_trades[:to_evict] for evict_trade in eviction_candidates: + evict_price = _price_for(evict_trade) pnl, pnl_pct = _calculate_pnl( - evict_trade.entry_price, current_price, evict_trade.direction, evict_trade.quantity + evict_trade.entry_price, evict_price, evict_trade.direction, evict_trade.quantity ) - evict_trade.exit_price = current_price + evict_trade.exit_price = evict_price evict_trade.exit_time = datetime.now(timezone.utc) evict_trade.exit_reason = "MAX_LIMIT_EVICT" evict_trade.pnl = pnl diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index 10964a3..c5c26e5 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -17,13 +17,19 @@ os.environ.setdefault("ENCRYPTION_KEY", "00" * 32) os.environ.setdefault("JWT_PRIVATE_KEY_PATH", "/nonexistent/jwt_private.pem") os.environ.setdefault("JWT_PUBLIC_KEYS_DIR", "/nonexistent/jwt_public_keys") +from datetime import timezone + import pytest import pytest_asyncio +from sqlalchemy import DateTime from sqlalchemy.dialects.postgresql import TIMESTAMP as PG_TIMESTAMP from sqlalchemy.dialects.postgresql import UUID as PG_UUID +from sqlalchemy.dialects.sqlite import DATETIME as SQLITE_DATETIME +from sqlalchemy.dialects.sqlite import dialect as _sqlite_dialect from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine from sqlalchemy.ext.compiler import compiles from sqlalchemy.pool import StaticPool +from sqlalchemy.types import TypeDecorator @compiles(PG_UUID, "sqlite") @@ -36,33 +42,70 @@ def _compile_pg_timestamp_sqlite(element, compiler, **kw): # noqa: ANN001 return "DATETIME" -@pytest_asyncio.fixture -async def db_session(): - """Fresh in-memory SQLite DB per test, with only the tables these tests need.""" - from app.database import Base - from app.models import AuditLog, Exchange, ExchangeCredential, HypotheticalTrade, Signal, User +class _UTCDateTime(TypeDecorator): + """SQLite has no native timezone-aware datetime storage — it silently + drops tzinfo on round-trip, which then blows up any `now - stored_value` + arithmetic in production code that (correctly) assumes tz-aware + datetimes throughout (as Postgres's TIMESTAMPTZ guarantees). Re-attach + UTC on the way out so model code doesn't need to know it's talking to + SQLite in tests.""" + + impl = SQLITE_DATETIME + cache_ok = True + + def process_bind_param(self, value, dialect): # noqa: ANN001 + if value is not None and value.tzinfo is not None: + return value.astimezone(timezone.utc).replace(tzinfo=None) + return value + + def process_result_value(self, value, dialect): # noqa: ANN001 + if value is not None and value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value + + +_sqlite_dialect.colspecs = { + **_sqlite_dialect.colspecs, + DateTime: _UTCDateTime, + PG_TIMESTAMP: _UTCDateTime, +} + + +def _needed_tables(): + from app.models import AuditLog, Exchange, ExchangeCredential, HypotheticalTrade, Signal, Symbol, User + from app.models.candle import Candle from app.models.real_trade import RealTrade + return [ + User.__table__, + Exchange.__table__, + ExchangeCredential.__table__, + RealTrade.__table__, + Signal.__table__, + HypotheticalTrade.__table__, + AuditLog.__table__, + Symbol.__table__, + Candle.__table__, + ] + + +async def _new_sqlite_engine(): engine = create_async_engine( "sqlite+aiosqlite:///:memory:", poolclass=StaticPool, connect_args={"check_same_thread": False}, ) + from app.database import Base async with engine.begin() as conn: - await conn.run_sync( - Base.metadata.create_all, - tables=[ - User.__table__, - Exchange.__table__, - ExchangeCredential.__table__, - RealTrade.__table__, - Signal.__table__, - HypotheticalTrade.__table__, - AuditLog.__table__, - ], - ) + await conn.run_sync(Base.metadata.create_all, tables=_needed_tables()) + return engine + +@pytest_asyncio.fixture +async def db_session(): + """Fresh in-memory SQLite DB per test, with only the tables these tests need.""" + engine = await _new_sqlite_engine() session = AsyncSession(engine, expire_on_commit=False) try: yield session @@ -71,5 +114,22 @@ async def db_session(): await engine.dispose() +@pytest_asyncio.fixture +async def session_factory(): + """Like `db_session`, but yields an `async_sessionmaker` instead of a + single session — for testing functions (e.g. in signal_service.py) that + open their own session(s) via a module-level `async_session_factory` + rather than accepting one as a parameter. Monkeypatch that module-level + name to this fixture's factory to redirect it at the in-memory DB.""" + from sqlalchemy.ext.asyncio import async_sessionmaker + + engine = await _new_sqlite_engine() + factory = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + try: + yield factory + finally: + await engine.dispose() + + def new_uuid() -> uuid.UUID: return uuid.uuid4() diff --git a/backend/tests/test_trade_executor.py b/backend/tests/test_trade_executor.py index 70a1ac5..fdb5bc6 100644 --- a/backend/tests/test_trade_executor.py +++ b/backend/tests/test_trade_executor.py @@ -20,7 +20,8 @@ from decimal import Decimal import pytest from sqlalchemy import select -from app.models import HypotheticalTrade, Signal, User +from app.models import Exchange, HypotheticalTrade, Signal, Symbol, User +from app.models.candle import Candle from app.services.trade_executor import ( MAX_OPEN_TRADES, STRONG_BUY, @@ -231,4 +232,67 @@ async def test_hybrid_eviction_evicts_oldest_when_all_open_trades_are_winners(db assert len(closed_trades) == 1 assert closed_trades[0].symbol == "SYM0/USDT", "FIFO: the oldest trade must be evicted when all are winners" assert closed_trades[0].exit_reason == "MAX_LIMIT_EVICT" - assert any(t.symbol == "NEW/USDT" for t in open_trades) + + +async def test_hybrid_eviction_uses_each_trades_own_symbol_price_not_incoming_signal_price(db_session): + """Regression test for finding (p): eviction PnL must be computed with + each open trade's OWN current price, not the price of whatever symbol + the currently-processed signal happens to be for. + + Setup: REAL/USDT is a real loser at its own market price (candle close + 50, entry 100 -> -50 PnL), but the incoming signal is for NEW/USDT at a + much higher price (1000). The old buggy code priced every open trade at + 1000, which would make REAL/USDT look like a huge winner (+900) and + evict the oldest FIFO trade instead of the actual loser. + """ + user = make_user() + exchange = Exchange(name="mexc", display_name="MEXC") + db_session.add_all([user, exchange]) + await db_session.flush() + + real_symbol = Symbol(symbol="REAL/USDT", base="REAL", quote="USDT", exchange_id=exchange.id) + db_session.add(real_symbol) + await db_session.flush() + + db_session.add(Candle( + symbol_id=real_symbol.id, timeframe="1h", timestamp=datetime.now(timezone.utc), + open=Decimal("60"), high=Decimal("65"), low=Decimal("48"), close=Decimal("50"), + volume=Decimal("1000"), + )) + + base_time = datetime.now(timezone.utc) - timedelta(days=1) + existing_trades = [ + HypotheticalTrade( + user_id=user.id, symbol="REAL/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=base_time, # oldest -> would be the FIFO pick if it were mis-priced as a winner + quantity=Decimal("1"), status="OPEN", + ) + ] + for i in range(MAX_OPEN_TRADES - 1): + existing_trades.append(HypotheticalTrade( + user_id=user.id, symbol=f"FILLER{i}/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=base_time + timedelta(minutes=i + 1), # all newer than REAL/USDT + quantity=Decimal("1"), status="OPEN", + )) + db_session.add_all(existing_trades) + await db_session.flush() + + signal = make_signal(STRONG_BUY, symbol="NEW/USDT") + db_session.add(signal) + await db_session.flush() + + await execute_signal_trade(db_session, signal, "NEW/USDT", "mexc", "1h", Decimal("1000")) + + all_trades = await _open_trades_for(db_session, user.id) + closed_trades = [t for t in all_trades if t.status == "CLOSED"] + + assert len(closed_trades) == 1 + assert closed_trades[0].symbol == "REAL/USDT", ( + "must evict the trade that is an actual loser at its OWN price, " + "not misjudge it as a winner using the unrelated incoming signal's price" + ) + assert closed_trades[0].exit_price == Decimal("50"), "exit price must come from REAL/USDT's own candle, not the signal's 1000" + assert closed_trades[0].pnl == Decimal("-50") + assert any(t.symbol == "NEW/USDT" and t.status == "OPEN" for t in all_trades)