fix: symbol-specific Kelly win rate, wire dead sizing code, safer walk-forward fallback

- signal_booster.py: _compute_rates() now also computes a per-symbol win
  rate (not just system-wide/direction aggregates); trade_executor.py's
  Kelly sizing prefers it when the symbol has enough closed-trade history.
  Added _as_datetime() to normalize closed_at across backends/drivers
  that return either a real datetime or a string from raw SQL.
- trade_executor.py: volatility-filter and Kelly-sizing exception handlers
  now log at warning level with the actual exception instead of silently
  swallowing failures that affect how much money a trade risks.
- risk_manager.py: compute_partial_tp_levels() now returns all 3 levels
  its docstring always promised (TP1 25% + TP2 35% + 40% trailing
  remainder) instead of silently dropping the last 40%.
- trade_executor.py: compute_volatility_adjusted_size() was dead code;
  now applied as a multiplier on the Kelly-derived trade_size (using
  max_risk_pct=100 to reinterpret it as "scale the already-sized trade"
  rather than "% of a bankroll", which would always collapse to this
  pipeline's $5 floor at its actual dollar scale).
- walk_forward.py: grid-search fallback (when every combo is too sparse
  to trust) now picks the combo with the most trades/highest PnL instead
  of always the grid's arbitrary first entry. Raised MIN_TRADES_PER_FOLD
  5 -> 15 for a more defensible statistical minimum.

219 backend tests pass (+10).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Le
2026-07-04 17:22:33 +07:00
parent cccef51eaf
commit 2399d0cb0c
9 changed files with 552 additions and 31 deletions
+14
View File
@@ -179,6 +179,12 @@ class AdaptiveSLTPOptimizer:
- TP1: ATR × 2.0 → close 25%
- TP2: ATR × 4.0 → close 35%
- Remainder: 40% with trailing stop
The three `close_percentage`s always sum to 1.0 (100% of the
position) — the first two are fixed-price take-profit levels, the
third has `price: None` and `trailing: True` since "trail the
remainder" is a runtime stop-management decision, not a price
this function can compute in isolation.
"""
params = REGIME_MULTIPLIERS.get(regime, REGIME_MULTIPLIERS["neutral"])
direction = direction.upper()
@@ -198,6 +204,14 @@ class AdaptiveSLTPOptimizer:
result.append({
"price": round(price, 8),
"close_percentage": level["close_pct"],
"trailing": False,
})
remainder_pct = 1.0 - sum(level["close_pct"] for level in levels)
result.append({
"price": None,
"close_percentage": round(remainder_pct, 8),
"trailing": True,
})
return result
+49 -6
View File
@@ -50,6 +50,31 @@ _STRATEGY_MAP: dict[str, str] = {
}
def _as_datetime(value: Any) -> datetime | None:
"""Normalize a raw `closed_at` value from `text()` SQL to a datetime.
Raw textual SQL (unlike ORM queries) carries no column type
information, so the driver returns whatever native type it stores
timestamps as — Postgres/asyncpg gives back a real `datetime`, but
SQLite/aiosqlite (used in this project's test suite) gives back a
plain string. Without this, the exponential-decay weighting below
would raise on any backend/driver that doesn't hand back a `datetime`.
"""
if value is None:
return None
if isinstance(value, str):
try:
value = datetime.fromisoformat(value)
except ValueError:
return None
if not isinstance(value, datetime):
return None
# This app stores/consumes timestamps as UTC throughout — a
# driver/backend that hands back a naive datetime (e.g. SQLite) means
# "UTC with the tzinfo stripped", not "some unspecified local time".
return value if value.tzinfo is not None else value.replace(tzinfo=timezone.utc)
# ── Core helpers ──────────────────────────────────────────────────────────
async def compute_strategy_win_rates(db: AsyncSession | None = None) -> dict[str, float]:
@@ -109,8 +134,8 @@ async def _compute_rates(db: AsyncSession) -> dict[str, float]:
for row in rows:
strategy = str(row[0])
is_win = int(row[1])
closed_at = row[2]
closed_at = _as_datetime(row[2])
if closed_at:
days_ago = (now_dt - closed_at).days
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
@@ -120,10 +145,13 @@ async def _compute_rates(db: AsyncSession) -> dict[str, float]:
strategy_weights[strategy] = strategy_weights.get(strategy, 0) + weight
strategy_wins[strategy] = strategy_wins.get(strategy, 0) + (weight if is_win else 0)
# Also compute direction-specific rates from the same data
# Also compute direction- and symbol-specific rates from the same
# data in one query (fix mm: Kelly sizing prefers a symbol's own
# win rate over the system-wide aggregate when there's enough
# history for that specific symbol).
dir_result = await db.execute(
text("""
SELECT direction,
SELECT direction, symbol,
CASE WHEN pnl > 0 THEN 1 ELSE 0 END,
closed_at
FROM hypothetical_trades
@@ -133,10 +161,13 @@ async def _compute_rates(db: AsyncSession) -> dict[str, float]:
LIMIT 5000
"""),
)
symbol_weights: dict[str, float] = {}
symbol_wins: dict[str, float] = {}
for row in dir_result.all():
direction = str(row[0]) if row[0] else "UNKNOWN"
is_win = int(row[1])
closed_at = row[2]
symbol = str(row[1]) if row[1] else None
is_win = int(row[2])
closed_at = _as_datetime(row[3])
if closed_at:
days_ago = (now_dt - closed_at).days
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
@@ -144,6 +175,9 @@ async def _compute_rates(db: AsyncSession) -> dict[str, float]:
weight = 0.5
direction_weights[direction] = direction_weights.get(direction, 0) + weight
direction_wins[direction] = direction_wins.get(direction, 0) + (weight if is_win else 0)
if symbol:
symbol_weights[symbol] = symbol_weights.get(symbol, 0) + weight
symbol_wins[symbol] = symbol_wins.get(symbol, 0) + (weight if is_win else 0)
rates: dict[str, float] = {}
total_weight = 0.0
@@ -172,6 +206,15 @@ async def _compute_rates(db: AsyncSession) -> dict[str, float]:
if dw >= MIN_TRADES * 0.5:
rates[f"__all___{direction}"] = ww / dw
# Symbol-specific rates (fix mm) — same MIN_TRADES bar as everything
# else here, so a thinly-traded symbol falls back to the direction/
# overall aggregate instead of a noisy few-trade estimate.
for symbol in symbol_weights:
sw = symbol_weights.get(symbol, 0)
sww = symbol_wins.get(symbol, 0)
if sw >= MIN_TRADES * 0.5:
rates[f"__symbol__{symbol}"] = sww / sw
_win_rate_cache = rates
_last_cache_update = datetime.now(timezone.utc)
await redis_client.set_json(_REDIS_KEY_WIN_RATES, rates, _CACHE_TTL_SECONDS)
+53 -5
View File
@@ -171,19 +171,29 @@ async def execute_signal_trade(
continue
# ── Volatility filter ──
# atr_pct_for_sizing survives past this block (fix oo) so the Kelly
# sizing step below can scale size down gracefully between the two
# skip thresholds, instead of ATR only ever being a binary
# skip/allow gate with no effect in between.
atr_pct_for_sizing: float | None = None
try:
snap = json.loads(signal.indicators_snapshot) if signal.indicators_snapshot else {}
atr_val = snap.get("atr_14")
if atr_val and isinstance(atr_val, list) and len(atr_val) > 0 and atr_val[-1]:
atr_pct = float(atr_val[-1]) / float(current_price) * 100
atr_pct_for_sizing = atr_pct
if atr_pct > 8.0:
logger.info("⛔ Skipping %s — ATR too high: %.2f%%", symbol, atr_pct)
continue
if atr_pct < 0.5:
logger.info("⛔ Skipping %s — ATR too low: %.2f%%", symbol, atr_pct)
continue
except Exception:
pass
except Exception as e:
# A parse/format error here must not silently disable the
# volatility gate in production with no trace of why — surface
# it, even though we still proceed (fail-open, matching the
# existing behavior of not blocking a trade over a data glitch).
logger.warning("Volatility filter failed for %s, proceeding without it: %s", symbol, e)
# ── Serialize per-user trade-opening decisions (fix ll) ──
# The `with_for_update()` on all_open_trades below only locks rows
@@ -288,12 +298,19 @@ async def execute_signal_trade(
kelly = DynamicKellySizer()
overall_rate = rates.get("__all__", 0.5)
dir_rate = rates.get(f"__all___{signal_direction}", overall_rate)
# (mm) Prefer this specific symbol's own win rate when there's
# enough history for it — a coin that trades very differently
# from the system-wide average (e.g. a consistently weaker
# altcoin) should be sized off its own edge, not the aggregate.
# Falls back to the direction/overall aggregate exactly as
# before when the symbol doesn't have enough closed trades yet.
win_rate = rates.get(f"__symbol__{symbol}", dir_rate)
signal_confidence = 0.5
if signal.indicators_snapshot:
snap = json.loads(signal.indicators_snapshot)
signal_confidence = snap.get("confidence", 0.5)
kelly_pct = kelly.compute_kelly_pct(
win_rate=dir_rate,
win_rate=win_rate,
avg_win=pnl_stats.get("avg_win", 3.0),
avg_loss=pnl_stats.get("avg_loss", 2.0),
confidence=signal_confidence,
@@ -313,8 +330,39 @@ async def execute_signal_trade(
kelly_pct *= 1.0 / math.sqrt(same_direction_open + 1)
if kelly_pct > 0:
trade_size = max(trade_size * Decimal(str(kelly_pct)), Decimal("1"))
except Exception:
logger.debug("Kelly sizing failed, using fixed trade_size")
# ── Volatility/regime adjustment (fix oo) ──
# `compute_volatility_adjusted_size` was previously never
# called anywhere — the volatility filter above only ever
# skipped a trade outright above/below its two hard cutoffs,
# with no graduated effect in between. `max_risk_pct=100` here
# means "scale 100% of the Kelly-derived trade_size by
# volatility/regime" rather than the function's own docstring
# framing ("% of a bankroll to risk") — that framing assumes a
# much larger base_size (an account balance) than this
# paper-trading pipeline's small fixed trade_size, where a
# literal 1-2% risk-per-trade would always collapse to the $5
# floor below regardless of volatility. Reusing the same
# vol_factor/regime_factor math as a pure multiplier on the
# already-sized trade instead keeps it meaningful at this
# pipeline's actual dollar scale.
if atr_pct_for_sizing is not None:
regime = "neutral"
if signal.indicators_snapshot:
regime = snap.get("market_regime") or "neutral"
trade_size = kelly.compute_volatility_adjusted_size(
base_size=trade_size,
atr_pct=Decimal(str(atr_pct_for_sizing)),
max_risk_pct=Decimal("100"),
regime=regime,
)
except Exception as e:
# Silently falling back here used to hide real bugs in the
# sizing pipeline (wrong rates shape, bad Decimal conversion,
# etc.) — this affects how much real/paper money a trade
# risks, so a failure here should be visible, not just a
# debug-level breadcrumb.
logger.warning("Kelly sizing failed for %s, using fixed trade_size: %s", symbol, e)
# Sane size bounds
trade_size = max(trade_size, Decimal("5"))
+32 -8
View File
@@ -47,7 +47,15 @@ DEFAULT_PARAM_GRID: dict[str, list[float]] = {
"max_hold_candles": [24, 48, 96],
}
MIN_TRADES_PER_FOLD = 5 # reject param combos too sparse to trust
# (fix rr) Reject param combos too sparse to trust. 5 trades is too thin a
# sample to estimate a Sharpe-like ratio's mean/std reliably — a couple of
# outlier trades can swing it wildly. 15 is still well short of the ~30
# quant practitioners often cite for a stable estimate, but demanding 30
# per fold would starve most folds of any "trustworthy" combo at all given
# this system's selective (STRONG-only) entry signals — 15 is a middle
# ground between statistical caution and having enough folds to walk
# forward over at all.
MIN_TRADES_PER_FOLD = 15
WARMUP_BUFFER_CANDLES = 60 # extra history fetched before each window so indicators aren't cold at window start
_TF_MINUTES = {"15m": 15, "30m": 30, "1h": 60, "4h": 240, "1d": 1440, "1w": 10080, "1M": 43200}
@@ -175,7 +183,20 @@ def _grid_search(
keys = list(param_grid.keys())
best_params: dict[str, float] | None = None
best_score = float("-inf")
best_stats: dict = {}
best_signals: list[dict] | None = None
best_trades: list[dict] | None = None
# (fix xx) If every combo is too sparse to trust (_fold_score returns
# -inf for all of them), we still need to report *something* for the
# fold — track the least-bad combo as we go instead of always falling
# back to the grid's arbitrary first entry, which could easily be the
# worst-performing one. Ranked by (trade count, total PnL): more trades
# means closer to being statistically trustworthy in the first place,
# and PnL breaks ties between equally-sparse combos.
fallback_params: dict[str, float] | None = None
fallback_rank: tuple[int, float] = (-1, float("-inf"))
fallback_signals: list[dict] | None = None
fallback_trades: list[dict] | None = None
for combo in product(*(param_grid[k] for k in keys)):
params = dict(zip(keys, combo))
@@ -185,15 +206,18 @@ def _grid_search(
if score > best_score:
best_score = score
best_params = params
best_stats = _compute_stats(all_signals, trades)
best_signals, best_trades = all_signals, trades
rank = (len(closed_trades), sum(float(t.get("pnl", 0.0)) for t in closed_trades))
if rank > fallback_rank:
fallback_rank = rank
fallback_params = params
fallback_signals, fallback_trades = all_signals, trades
if best_params is None:
# Every combo scored -inf (too few trades) — still report the
# grid's first combination so the fold has *something* to show.
best_params = {k: param_grid[k][0] for k in keys}
all_signals, trades = _run_combo(candles, scores_series, trade_size, best_params, fee_pct, slippage_pct)
best_stats = _compute_stats(all_signals, trades)
best_params, best_signals, best_trades = fallback_params, fallback_signals, fallback_trades
best_stats = _compute_stats(best_signals, best_trades)
return best_params, best_score, best_stats
+19 -5
View File
@@ -134,21 +134,35 @@ class TestAdaptiveSLTPOptimizerComputeSlTp:
class TestAdaptiveSLTPOptimizerPartialTpLevels:
"""Regression tests for fix (nn): the docstring always promised TP1
25% + TP2 35% + a 40% trailing remainder (100% of the position
accounted for), but the code only ever returned the first two levels
(60% total) — silently leaving the other 40% unaccounted for from a
caller's point of view.
"""
def setup_method(self):
self.opt = AdaptiveSLTPOptimizer()
def test_returns_two_levels_summing_close_percentage_below_one(self):
def test_returns_three_levels_summing_close_percentage_to_one(self):
levels = self.opt.compute_partial_tp_levels(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
assert len(levels) == 2
assert len(levels) == 3
total_close_pct = sum(lvl["close_percentage"] for lvl in levels)
assert 0 < total_close_pct <= 1.0
assert total_close_pct == pytest.approx(1.0)
def test_first_two_levels_have_fixed_prices_third_is_trailing_remainder(self):
levels = self.opt.compute_partial_tp_levels(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
assert levels[0]["price"] is not None and levels[0]["trailing"] is False
assert levels[1]["price"] is not None and levels[1]["trailing"] is False
assert levels[2]["price"] is None and levels[2]["trailing"] is True
assert levels[2]["close_percentage"] == pytest.approx(0.40)
def test_long_levels_are_above_entry_and_increasing(self):
levels = self.opt.compute_partial_tp_levels(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
assert levels[0]["price"] < levels[1]["price"]
assert all(lvl["price"] > 50_000.0 for lvl in levels)
assert all(lvl["price"] > 50_000.0 for lvl in levels[:2])
def test_short_levels_are_below_entry_and_decreasing(self):
levels = self.opt.compute_partial_tp_levels(atr=100.0, entry_price=50_000.0, regime="trending", direction="SHORT")
assert levels[0]["price"] > levels[1]["price"]
assert all(lvl["price"] < 50_000.0 for lvl in levels)
assert all(lvl["price"] < 50_000.0 for lvl in levels[:2])
+79
View File
@@ -0,0 +1,79 @@
"""Tests for app/services/signal_booster.py's `_compute_rates` — the raw-SQL
win-rate aggregation that Kelly sizing reads from. Runs against a real
(in-memory SQLite) DB so the actual query logic is exercised, not mocked.
"""
from __future__ import annotations
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from app.models.signal import HypotheticalTrade
from app.services.signal_booster import _compute_rates
def make_trade(symbol, direction, pnl, entry_reason="double_bb_rsi", closed_at=None):
now = datetime.now(timezone.utc)
return HypotheticalTrade(
symbol=symbol, exchange="mexc", timeframe="1h", direction=direction,
entry_price=Decimal("100"), entry_time=now - timedelta(hours=1),
entry_reason=entry_reason, exit_price=Decimal("110" if pnl > 0 else "90"),
exit_time=now, quantity=Decimal("1"), pnl=Decimal(str(pnl)),
pnl_percent=Decimal("10"), status="CLOSED",
closed_at=closed_at or now,
)
class TestSymbolSpecificRates:
"""Regression tests for fix (mm): Kelly sizing used to only ever see a
system-wide `__all__` (and direction-level `__all___{LONG,SHORT}`) win
rate — never a specific symbol's own performance, even when that
symbol has plenty of its own closed-trade history that trades very
differently from the aggregate.
"""
async def test_symbol_with_enough_trades_gets_its_own_rate(self, db_session):
# BTC/USDT: 8 wins, 0 losses (needs >= MIN_TRADES*0.5 = 7.5 weighted trades).
for _ in range(8):
db_session.add(make_trade("BTC/USDT", "LONG", pnl=10.0))
await db_session.flush()
rates = await _compute_rates(db_session)
assert rates["__symbol__BTC/USDT"] == 1.0
async def test_symbol_with_too_few_trades_has_no_own_rate(self, db_session):
# Only 2 trades — well under the MIN_TRADES*0.5 threshold.
db_session.add(make_trade("ETH/USDT", "LONG", pnl=10.0))
db_session.add(make_trade("ETH/USDT", "LONG", pnl=-10.0))
await db_session.flush()
rates = await _compute_rates(db_session)
assert "__symbol__ETH/USDT" not in rates
async def test_different_symbols_do_not_share_a_rate(self, db_session):
for _ in range(8):
db_session.add(make_trade("WINNER/USDT", "LONG", pnl=10.0))
for _ in range(8):
db_session.add(make_trade("LOSER/USDT", "LONG", pnl=-10.0))
await db_session.flush()
rates = await _compute_rates(db_session)
assert rates["__symbol__WINNER/USDT"] == 1.0
assert rates["__symbol__LOSER/USDT"] == 0.0
async def test_direction_and_overall_rates_are_still_computed(self, db_session):
"""The symbol-rate addition must not break the existing __all__ /
__all___{direction} aggregates it was computed alongside."""
for _ in range(8):
db_session.add(make_trade("BTC/USDT", "LONG", pnl=10.0))
for _ in range(8):
db_session.add(make_trade("BTC/USDT", "SHORT", pnl=-10.0))
await db_session.flush()
rates = await _compute_rates(db_session)
assert rates["__all___LONG"] == 1.0
assert rates["__all___SHORT"] == 0.0
assert rates["__all__"] == 0.5
+136
View File
@@ -450,6 +450,142 @@ class TestKellyPortfolioCorrelationDampening:
assert float(hedged_trade.quantity) == pytest.approx(float(baseline_trade.quantity), rel=1e-6)
class TestKellySymbolSpecificWinRate:
"""Regression test for fix (mm): Kelly sizing used to always size off
the system-wide `__all___{direction}` win rate, even when the specific
symbol being traded has its own (very different) win-rate history with
plenty of samples — now it prefers the symbol's own rate when present.
"""
async def test_symbol_specific_rate_overrides_direction_aggregate(self, db_session, monkeypatch):
import app.services.signal_booster as signal_booster_module
async def fake_rates():
return {
"__all__": 0.5, "__all___LONG": 0.5,
# BTC/USDT trades far better than the system-wide average.
"__symbol__BTC/USDT": 0.9,
}
async def fake_pnl_stats():
return {"avg_win": 5.0, "avg_loss": 1.0}
monkeypatch.setattr(signal_booster_module, "get_cached_rates", fake_rates)
monkeypatch.setattr(signal_booster_module, "get_pnl_stats", fake_pnl_stats)
# Restrict each user to only their own symbol — otherwise both
# users (empty auto_trade_tokens = trades everything) would get
# BOTH the ETH and BTC trade from each call below, and the second
# user's open-position count from the first call would trigger
# portfolio-correlation dampening (fix hh) that has nothing to do
# with what this test is isolating.
weak_direction_user = make_user(trade_size=200, auto_trade_tokens=["ETH/USDT"])
strong_symbol_user = make_user(trade_size=200, auto_trade_tokens=["BTC/USDT"])
db_session.add_all([weak_direction_user, strong_symbol_user])
await db_session.flush()
confident_signal_eth = Signal(
symbol="ETH/USDT", exchange="mexc", timeframe="1h",
signal_type=STRONG_BUY, strength="STRONG", price=Decimal("100"),
timestamp=datetime.now(timezone.utc),
indicators_snapshot=json.dumps({"confidence": 1.0}),
)
confident_signal_btc = Signal(
symbol="BTC/USDT", exchange="mexc", timeframe="1h",
signal_type=STRONG_BUY, strength="STRONG", price=Decimal("100"),
timestamp=datetime.now(timezone.utc),
indicators_snapshot=json.dumps({"confidence": 1.0}),
)
db_session.add_all([confident_signal_eth, confident_signal_btc])
await db_session.flush()
# ETH/USDT has no rate of its own -> falls back to __all___LONG (0.5).
await execute_signal_trade(db_session, confident_signal_eth, "ETH/USDT", "mexc", "1h", Decimal("50000"))
# BTC/USDT has its own, much higher, rate (0.9) -> should size larger.
await execute_signal_trade(db_session, confident_signal_btc, "BTC/USDT", "mexc", "1h", Decimal("50000"))
eth_trade = (await _open_trades_for(db_session, weak_direction_user.id, "ETH/USDT"))[0]
btc_trade = (await _open_trades_for(db_session, strong_symbol_user.id, "BTC/USDT"))[0]
assert float(btc_trade.quantity) > float(eth_trade.quantity)
class TestVolatilityRegimeSizeAdjustment:
"""Regression test for fix (oo): `compute_volatility_adjusted_size`
was dead code — the volatility filter only ever skipped a trade
outright above 8%/below 0.5% ATR, with no effect at all in between.
It's now applied as a multiplier on the Kelly-derived trade_size using
the same ATR%/regime already read off the signal snapshot.
"""
def _signal_with_atr_and_regime(self, symbol, atr_abs, regime):
return Signal(
symbol=symbol, exchange="mexc", timeframe="1h",
signal_type=STRONG_BUY, strength="STRONG", price=Decimal("100"),
timestamp=datetime.now(timezone.utc),
indicators_snapshot=json.dumps({
"atr_14": [atr_abs], "confidence": 1.0, "market_regime": regime,
}),
)
def _patch_kelly_inputs(self, monkeypatch):
import app.services.signal_booster as signal_booster_module
async def fake_rates():
return {"__all__": 0.6, "__all___LONG": 0.6}
async def fake_pnl_stats():
return {"avg_win": 3.0, "avg_loss": 2.0}
monkeypatch.setattr(signal_booster_module, "get_cached_rates", fake_rates)
monkeypatch.setattr(signal_booster_module, "get_pnl_stats", fake_pnl_stats)
async def test_volatile_regime_and_high_atr_shrinks_size_vs_calm_neutral(self, db_session, monkeypatch):
self._patch_kelly_inputs(monkeypatch)
calm_user = make_user(trade_size=200, auto_trade_tokens=["CALM/USDT"])
volatile_user = make_user(trade_size=200, auto_trade_tokens=["WILD/USDT"])
db_session.add_all([calm_user, volatile_user])
await db_session.flush()
# current_price=100 -> atr_pct = atr_abs (since atr_abs/100*100 = atr_abs).
calm_signal = self._signal_with_atr_and_regime("CALM/USDT", atr_abs=2.0, regime="neutral")
volatile_signal = self._signal_with_atr_and_regime("WILD/USDT", atr_abs=6.0, regime="volatile")
db_session.add_all([calm_signal, volatile_signal])
await db_session.flush()
await execute_signal_trade(db_session, calm_signal, "CALM/USDT", "mexc", "1h", Decimal("100"))
await execute_signal_trade(db_session, volatile_signal, "WILD/USDT", "mexc", "1h", Decimal("100"))
calm_trade = (await _open_trades_for(db_session, calm_user.id, "CALM/USDT"))[0]
volatile_trade = (await _open_trades_for(db_session, volatile_user.id, "WILD/USDT"))[0]
assert float(volatile_trade.quantity) < float(calm_trade.quantity)
async def test_no_atr_data_leaves_sizing_unaffected(self, db_session, monkeypatch):
"""No atr_14 in the snapshot (e.g. an older/partial signal) must
skip this adjustment entirely rather than erroring or applying a
default that changes existing sizing behavior."""
self._patch_kelly_inputs(monkeypatch)
user = make_user(trade_size=200)
db_session.add(user)
await db_session.flush()
signal = Signal(
symbol="BTC/USDT", exchange="mexc", timeframe="1h",
signal_type=STRONG_BUY, strength="STRONG", price=Decimal("100"),
timestamp=datetime.now(timezone.utc),
indicators_snapshot=json.dumps({"confidence": 1.0}), # no atr_14
)
db_session.add(signal)
await db_session.flush()
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("100"))
trade = (await _open_trades_for(db_session, user.id, "BTC/USDT"))[0]
assert trade.quantity > 0
def make_real_trade(user_id, symbol="BTC/USDT", side="buy", amount="1", price="100",
filled_amount=None, status="filled", created_at=None) -> RealTrade:
return RealTrade(
+56 -7
View File
@@ -51,10 +51,11 @@ def test_fold_score_rejects_too_few_trades():
def test_fold_score_prefers_consistent_edge_over_lucky_streak():
# Same total PnL (50), but one is steady small wins, the other is one
# Same total PnL (150), but one is steady small wins, the other is one
# huge win plus several losses — the steadier one should score higher.
consistent = [{"pnl": 10.0} for _ in range(5)]
lucky = [{"pnl": 50.0}, {"pnl": -10.0}, {"pnl": -10.0}, {"pnl": -10.0}, {"pnl": -10.0}]
# Both need >= MIN_TRADES_PER_FOLD trades to get a real (non -inf) score.
consistent = [{"pnl": 10.0} for _ in range(15)]
lucky = [{"pnl": 290.0}] + [{"pnl": -10.0} for _ in range(14)]
score_consistent = walk_forward._fold_score(consistent)
score_lucky = walk_forward._fold_score(lucky)
@@ -63,9 +64,9 @@ def test_fold_score_prefers_consistent_edge_over_lucky_streak():
def test_fold_score_zero_variance_all_same_sign():
trades = [{"pnl": 10.0} for _ in range(6)]
trades = [{"pnl": 10.0} for _ in range(16)]
score = walk_forward._fold_score(trades)
assert score == pytest.approx(10.0 * math.sqrt(6))
assert score == pytest.approx(10.0 * math.sqrt(16))
# ── _max_drawdown_pct ──────────────────────────────────────────────────
@@ -91,7 +92,7 @@ async def test_grid_search_picks_the_best_scoring_combo(monkeypatch):
def fake_run_combo(candles, scores_series, trade_size, params, fee_pct=None, slippage_pct=None):
if params == good_params:
trades = [{"pnl": 10.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 110.0} for _ in range(10)]
trades = [{"pnl": 10.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 110.0} for _ in range(15)]
else:
trades = [
{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0},
@@ -107,7 +108,7 @@ async def test_grid_search_picks_the_best_scoring_combo(monkeypatch):
)
assert best_params == good_params
assert best_stats["trades"]["total"] == 10
assert best_stats["trades"]["total"] == 15
assert best_score > float("-inf")
@@ -130,6 +131,54 @@ async def test_grid_search_falls_back_when_every_combo_too_sparse(monkeypatch):
assert best_stats["trades"]["total"] == 1
@pytest.mark.asyncio
async def test_grid_search_fallback_prefers_more_trades_not_first_combo(monkeypatch):
"""Regression test for fix (xx): when every combo is too sparse to
trust via _fold_score, the fallback used to always pick the grid's
first entry regardless of how sparse/lucky it was. It should now pick
whichever sparse combo has the most trades (closer to statistically
meaningful) — here, the SECOND combo (4.5) — not the first (3.5)."""
def fake_run_combo(candles, scores_series, trade_size, params, fee_pct=None, slippage_pct=None):
if params["strong_threshold"] == 3.5:
trades = [{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0}]
else:
trades = [
{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0},
{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0},
{"pnl": 1.0, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0},
]
return [], trades
monkeypatch.setattr(walk_forward, "_run_combo", fake_run_combo)
grid = {"strong_threshold": [3.5, 4.5], "signal_threshold": [1.0], "max_hold_candles": [48]}
best_params, best_score, best_stats = walk_forward._grid_search(
candles=[], scores_series=[], param_grid=grid, trade_size=Decimal("10"),
)
assert best_params == {"strong_threshold": 4.5, "signal_threshold": 1.0, "max_hold_candles": 48}
assert best_score == float("-inf")
assert best_stats["trades"]["total"] == 3
@pytest.mark.asyncio
async def test_grid_search_fallback_breaks_ties_with_higher_pnl(monkeypatch):
"""Same trade count for both sparse combos -> tie-break on total PnL."""
def fake_run_combo(candles, scores_series, trade_size, params, fee_pct=None, slippage_pct=None):
pnl = 1.0 if params["strong_threshold"] == 3.5 else 5.0
return [], [{"pnl": pnl, "status": "CLOSED", "entry_price": 100.0, "exit_price": 101.0}]
monkeypatch.setattr(walk_forward, "_run_combo", fake_run_combo)
grid = {"strong_threshold": [3.5, 4.5], "signal_threshold": [1.0], "max_hold_candles": [48]}
best_params, _best_score, best_stats = walk_forward._grid_search(
candles=[], scores_series=[], param_grid=grid, trade_size=Decimal("10"),
)
assert best_params == {"strong_threshold": 4.5, "signal_threshold": 1.0, "max_hold_candles": 48}
assert best_stats["trades"]["total_pnl"] == pytest.approx(5.0)
# ── end-to-end (small synthetic dataset) ─────────────────────────────────
async def _seed_symbol(db_session, name="BTC/USDT", exchange_name="mexc"):