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:
@@ -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
@@ -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,23 +42,41 @@ 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
|
||||
|
||||
engine = create_async_engine(
|
||||
"sqlite+aiosqlite:///:memory:",
|
||||
poolclass=StaticPool,
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(
|
||||
Base.metadata.create_all,
|
||||
tables=[
|
||||
return [
|
||||
User.__table__,
|
||||
Exchange.__table__,
|
||||
ExchangeCredential.__table__,
|
||||
@@ -60,9 +84,28 @@ async def db_session():
|
||||
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=_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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user