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.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
+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_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()
+66 -2
View File
@@ -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)