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 <noreply@anthropic.com>
This commit is contained in:
Le
2026-07-03 22:29:41 +07:00
parent 4ab2dfbe7c
commit 9a0d2ab220
3 changed files with 183 additions and 22 deletions
+40 -3
View File
@@ -15,8 +15,11 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.core.exceptions import AppException from app.core.exceptions import AppException
from app.database import async_session_factory 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.real_trade import RealTrade
from app.models.signal import HypotheticalTrade, Signal from app.models.signal import HypotheticalTrade, Signal
from app.models.symbol import Symbol
from app.models.user import User from app.models.user import User
from app.services.audit_service import log_action from app.services.audit_service import log_action
@@ -195,9 +198,42 @@ async def execute_signal_trade(
if open_count >= MAX_OPEN_TRADES: if open_count >= MAX_OPEN_TRADES:
to_evict = open_count - MAX_OPEN_TRADES + 1 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 = [] open_with_pnl = []
for t in all_open_trades: 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)) open_with_pnl.append((t, pnl_val))
losers = [(t, pnl) for t, pnl in open_with_pnl if pnl < 0] 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] eviction_candidates = all_open_trades[:to_evict]
for evict_trade in eviction_candidates: for evict_trade in eviction_candidates:
evict_price = _price_for(evict_trade)
pnl, pnl_pct = _calculate_pnl( 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_time = datetime.now(timezone.utc)
evict_trade.exit_reason = "MAX_LIMIT_EVICT" evict_trade.exit_reason = "MAX_LIMIT_EVICT"
evict_trade.pnl = pnl evict_trade.pnl = pnl
+77 -17
View File
@@ -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_PRIVATE_KEY_PATH", "/nonexistent/jwt_private.pem")
os.environ.setdefault("JWT_PUBLIC_KEYS_DIR", "/nonexistent/jwt_public_keys") os.environ.setdefault("JWT_PUBLIC_KEYS_DIR", "/nonexistent/jwt_public_keys")
from datetime import timezone
import pytest import pytest
import pytest_asyncio import pytest_asyncio
from sqlalchemy import DateTime
from sqlalchemy.dialects.postgresql import TIMESTAMP as PG_TIMESTAMP from sqlalchemy.dialects.postgresql import TIMESTAMP as PG_TIMESTAMP
from sqlalchemy.dialects.postgresql import UUID as PG_UUID 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.asyncio import AsyncSession, create_async_engine
from sqlalchemy.ext.compiler import compiles from sqlalchemy.ext.compiler import compiles
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from sqlalchemy.types import TypeDecorator
@compiles(PG_UUID, "sqlite") @compiles(PG_UUID, "sqlite")
@@ -36,33 +42,70 @@ def _compile_pg_timestamp_sqlite(element, compiler, **kw): # noqa: ANN001
return "DATETIME" return "DATETIME"
@pytest_asyncio.fixture class _UTCDateTime(TypeDecorator):
async def db_session(): """SQLite has no native timezone-aware datetime storage — it silently
"""Fresh in-memory SQLite DB per test, with only the tables these tests need.""" drops tzinfo on round-trip, which then blows up any `now - stored_value`
from app.database import Base arithmetic in production code that (correctly) assumes tz-aware
from app.models import AuditLog, Exchange, ExchangeCredential, HypotheticalTrade, Signal, User 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 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( engine = create_async_engine(
"sqlite+aiosqlite:///:memory:", "sqlite+aiosqlite:///:memory:",
poolclass=StaticPool, poolclass=StaticPool,
connect_args={"check_same_thread": False}, connect_args={"check_same_thread": False},
) )
from app.database import Base
async with engine.begin() as conn: async with engine.begin() as conn:
await conn.run_sync( await conn.run_sync(Base.metadata.create_all, tables=_needed_tables())
Base.metadata.create_all, return engine
tables=[
User.__table__,
Exchange.__table__,
ExchangeCredential.__table__,
RealTrade.__table__,
Signal.__table__,
HypotheticalTrade.__table__,
AuditLog.__table__,
],
)
@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) session = AsyncSession(engine, expire_on_commit=False)
try: try:
yield session yield session
@@ -71,5 +114,22 @@ async def db_session():
await engine.dispose() 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: def new_uuid() -> uuid.UUID:
return uuid.uuid4() return uuid.uuid4()
+66 -2
View File
@@ -20,7 +20,8 @@ from decimal import Decimal
import pytest import pytest
from sqlalchemy import select 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 ( from app.services.trade_executor import (
MAX_OPEN_TRADES, MAX_OPEN_TRADES,
STRONG_BUY, 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 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].symbol == "SYM0/USDT", "FIFO: the oldest trade must be evicted when all are winners"
assert closed_trades[0].exit_reason == "MAX_LIMIT_EVICT" 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)