test: them 81 pytest cho auth/orders/security/CORS/risk_manager/trade_executor/signal_service + CI
Backend truoc day chi co script goi httpx vao server dang chay that (test_auth.py, test_full_api.py), khong phai pytest that. Them bo test chay doc lap bang SQLite in-memory (khong can Postgres/Docker): - test_rbac_deps.py: RBAC chain + regression-guard cho fix vai tro o /orders/place - test_order_exchange_routing.py: routing dung san theo credential - test_security_encryption.py: AES-GCM round-trip + tuong thich nguoc AES-CBC - test_cors_config.py: CORS fail-closed khi thieu cau hinh - test_risk_manager.py: Kelly sizing + SL/TP adaptive theo tung regime - test_trade_executor.py: STRONG-only, dedup, reversal, volatility filter, hybrid eviction FIFO -- toan bo quy tac mo/dong trade - test_signal_service_scoring.py: he thong cham diem 13 thuat toan Them .gitea/workflows/backend-tests.yml chay pytest tu dong khi push/PR dung vao backend/** (can Gitea Actions + runner da duoc bat tren instance). Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
[pytest]
|
||||
asyncio_mode = auto
|
||||
testpaths = tests
|
||||
@@ -0,0 +1,5 @@
|
||||
# Dev/test-only dependencies. Install on top of requirements.txt:
|
||||
# pip install -r requirements.txt -r requirements-dev.txt
|
||||
pytest==9.1.1
|
||||
pytest-asyncio==1.4.0
|
||||
aiosqlite==0.22.1
|
||||
@@ -0,0 +1,75 @@
|
||||
"""Shared pytest fixtures for backend tests.
|
||||
|
||||
These tests run against an in-memory SQLite database instead of the real
|
||||
Postgres instance. A couple of Postgres-only column types used by the ORM
|
||||
models (UUID, TIMESTAMP) are given SQLite-compatible renderings via
|
||||
``@compiles`` so the *specific* tables under test can be created without
|
||||
touching production model code.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import uuid
|
||||
|
||||
# Must be set before `app.config` (and anything importing it) loads, so the
|
||||
# encryption round-trip in security.py has a valid key to work with.
|
||||
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")
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy.dialects.postgresql import TIMESTAMP as PG_TIMESTAMP
|
||||
from sqlalchemy.dialects.postgresql import UUID as PG_UUID
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
|
||||
from sqlalchemy.ext.compiler import compiles
|
||||
from sqlalchemy.pool import StaticPool
|
||||
|
||||
|
||||
@compiles(PG_UUID, "sqlite")
|
||||
def _compile_pg_uuid_sqlite(element, compiler, **kw): # noqa: ANN001
|
||||
return "CHAR(32)"
|
||||
|
||||
|
||||
@compiles(PG_TIMESTAMP, "sqlite")
|
||||
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
|
||||
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=[
|
||||
User.__table__,
|
||||
Exchange.__table__,
|
||||
ExchangeCredential.__table__,
|
||||
RealTrade.__table__,
|
||||
Signal.__table__,
|
||||
HypotheticalTrade.__table__,
|
||||
AuditLog.__table__,
|
||||
],
|
||||
)
|
||||
|
||||
session = AsyncSession(engine, expire_on_commit=False)
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
await session.close()
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
def new_uuid() -> uuid.UUID:
|
||||
return uuid.uuid4()
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Tests for fix (g): an unset CORS_ORIGINS must fail closed (deny all
|
||||
cross-origin requests), never fall back to the wildcard "*" — which combined
|
||||
with allow_credentials=True is a real vulnerability, not just a lint nit.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _reload_main_api_with_cors(monkeypatch, cors_origins: str):
|
||||
monkeypatch.setenv("CORS_ORIGINS", cors_origins)
|
||||
from app import config
|
||||
|
||||
importlib.reload(config)
|
||||
monkeypatch.setattr("app.config.settings", config.Settings())
|
||||
|
||||
from app import main_api
|
||||
|
||||
importlib.reload(main_api)
|
||||
return main_api
|
||||
|
||||
|
||||
def test_empty_cors_origins_denies_all_by_default(monkeypatch):
|
||||
main_api = _reload_main_api_with_cors(monkeypatch, "")
|
||||
assert main_api.origins == [], "must deny all cross-origin requests, not fall back to '*'"
|
||||
|
||||
|
||||
def test_configured_cors_origins_are_parsed(monkeypatch):
|
||||
main_api = _reload_main_api_with_cors(monkeypatch, "https://trading.dangloica.org, https://admin.dangloica.org")
|
||||
assert main_api.origins == ["https://trading.dangloica.org", "https://admin.dangloica.org"]
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _restore_main_api_after_test():
|
||||
"""Reload main_api once more after each test so later test modules that
|
||||
import it don't see a module left in a monkeypatched state."""
|
||||
yield
|
||||
import importlib as _importlib
|
||||
|
||||
from app import config as _config
|
||||
|
||||
_importlib.reload(_config)
|
||||
from app import main_api as _main_api
|
||||
|
||||
_importlib.reload(_main_api)
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Tests for fix (c): POST /orders/place must route to the exchange the user
|
||||
actually asked for (or their only/most-recent active credential), instead of
|
||||
being hardcoded to "mexc" regardless of the user's connected exchanges.
|
||||
|
||||
Runs against a real (in-memory SQLite) DB so the SQL filtering logic in
|
||||
`place_order` is genuinely exercised, not just mocked away.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.api.v1 import orders as orders_module
|
||||
from app.core.exceptions import NotFoundException
|
||||
from app.core.security import encrypt_api_key
|
||||
from app.exchange.types import OrderRequest
|
||||
from app.models import Exchange, ExchangeCredential, User
|
||||
from app.models.real_trade import RealTrade
|
||||
|
||||
|
||||
class StubAdapter:
|
||||
"""Replaces the real CCXT-backed exchange adapter in tests."""
|
||||
|
||||
def __init__(self):
|
||||
self.orders_placed: list[OrderRequest] = []
|
||||
|
||||
async def create_order(self, req: OrderRequest):
|
||||
self.orders_placed.append(req)
|
||||
return SimpleNamespace(
|
||||
order_id="STUB-ORDER-1",
|
||||
filled=req.amount,
|
||||
status="closed",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def trader_user():
|
||||
return User(
|
||||
id=uuid.uuid4(),
|
||||
username="trader1",
|
||||
email="trader1@example.com",
|
||||
password_hash="x",
|
||||
role="trader",
|
||||
is_active=True,
|
||||
)
|
||||
|
||||
|
||||
async def _seed(db_session, user: User):
|
||||
"""Create mexc + binance exchanges, and one credential each for `user`.
|
||||
|
||||
The mexc credential is created "earlier" than the binance one so a test
|
||||
can assert the no-exchange-given fallback picks the most recent (binance).
|
||||
"""
|
||||
mexc = Exchange(name="mexc", display_name="MEXC")
|
||||
binance = Exchange(name="binance", display_name="Binance")
|
||||
db_session.add_all([mexc, user])
|
||||
await db_session.flush()
|
||||
|
||||
older = datetime.now(timezone.utc) - timedelta(hours=1)
|
||||
newer = datetime.now(timezone.utc)
|
||||
|
||||
secret_enc, iv = encrypt_api_key("mexc-secret")
|
||||
mexc_cred = ExchangeCredential(
|
||||
id=uuid.uuid4(),
|
||||
user_id=user.id,
|
||||
exchange_id=mexc.id,
|
||||
api_key="mexc-key",
|
||||
api_secret_enc=secret_enc,
|
||||
api_secret_iv=iv,
|
||||
is_active=True,
|
||||
created_at=older,
|
||||
)
|
||||
db_session.add_all([binance, mexc_cred])
|
||||
await db_session.flush()
|
||||
|
||||
secret_enc2, iv2 = encrypt_api_key("binance-secret")
|
||||
binance_cred = ExchangeCredential(
|
||||
id=uuid.uuid4(),
|
||||
user_id=user.id,
|
||||
exchange_id=binance.id,
|
||||
api_key="binance-key",
|
||||
api_secret_enc=secret_enc2,
|
||||
api_secret_iv=iv2,
|
||||
is_active=True,
|
||||
created_at=newer,
|
||||
)
|
||||
db_session.add(binance_cred)
|
||||
await db_session.flush()
|
||||
return mexc, binance
|
||||
|
||||
|
||||
def _make_request(exchange: str | None) -> OrderRequest:
|
||||
return OrderRequest(
|
||||
symbol="BTC/USDT",
|
||||
side="buy",
|
||||
order_type="market",
|
||||
amount=Decimal("10"),
|
||||
exchange=exchange,
|
||||
)
|
||||
|
||||
|
||||
async def test_place_order_uses_requested_exchange(db_session, trader_user, monkeypatch):
|
||||
await _seed(db_session, trader_user)
|
||||
|
||||
calls: list[str] = []
|
||||
stub = StubAdapter()
|
||||
|
||||
def fake_create(name, api_key="", api_secret="", testnet=False):
|
||||
calls.append(name)
|
||||
return stub
|
||||
|
||||
monkeypatch.setattr(orders_module.exchange_factory, "create", fake_create)
|
||||
|
||||
req = _make_request(exchange="binance")
|
||||
order = await orders_module.place_order(req, db=db_session, current_user=trader_user)
|
||||
|
||||
assert calls == ["binance"], "must create the adapter for the exchange the caller asked for"
|
||||
assert order.order_id == "STUB-ORDER-1"
|
||||
|
||||
result = await db_session.execute(select(RealTrade).where(RealTrade.user_id == trader_user.id))
|
||||
trade = result.scalars().one()
|
||||
assert trade.exchange == "binance"
|
||||
|
||||
|
||||
async def test_place_order_404_when_no_credential_for_requested_exchange(db_session, trader_user, monkeypatch):
|
||||
await _seed(db_session, trader_user)
|
||||
monkeypatch.setattr(orders_module.exchange_factory, "create", lambda *a, **kw: StubAdapter())
|
||||
|
||||
req = _make_request(exchange="bybit") # user has no bybit credential
|
||||
with pytest.raises(NotFoundException) as excinfo:
|
||||
await orders_module.place_order(req, db=db_session, current_user=trader_user)
|
||||
assert "bybit" in excinfo.value.detail
|
||||
|
||||
|
||||
async def test_place_order_falls_back_to_most_recent_credential_when_exchange_omitted(
|
||||
db_session, trader_user, monkeypatch
|
||||
):
|
||||
await _seed(db_session, trader_user) # mexc created 1h before binance
|
||||
|
||||
calls: list[str] = []
|
||||
monkeypatch.setattr(
|
||||
orders_module.exchange_factory,
|
||||
"create",
|
||||
lambda name, **kw: (calls.append(name), StubAdapter())[1],
|
||||
)
|
||||
|
||||
req = _make_request(exchange=None)
|
||||
await orders_module.place_order(req, db=db_session, current_user=trader_user)
|
||||
|
||||
assert calls == ["binance"], "with no exchange specified, must fall back to the most recently added credential"
|
||||
@@ -0,0 +1,78 @@
|
||||
"""Tests for the RBAC dependency chain in app/core/deps.py.
|
||||
|
||||
Also guards against regressing fix (b): POST /orders/place must require the
|
||||
`trader` (or `admin`) role, not just "any authenticated user" — a `viewer`
|
||||
account must never be able to place a real order.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
from app.core import deps
|
||||
from app.core.exceptions import ForbiddenException, InvalidCredentialsException
|
||||
|
||||
|
||||
def fake_user(*, is_active=True, role="trader", is_admin=False):
|
||||
return SimpleNamespace(is_active=is_active, role=role, is_admin=is_admin, username="u")
|
||||
|
||||
|
||||
async def test_get_current_active_user_rejects_inactive():
|
||||
with pytest.raises(InvalidCredentialsException):
|
||||
await deps.get_current_active_user(current_user=fake_user(is_active=False))
|
||||
|
||||
|
||||
async def test_get_current_active_user_accepts_active():
|
||||
user = fake_user(is_active=True)
|
||||
result = await deps.get_current_active_user(current_user=user)
|
||||
assert result is user
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", ["trader", "admin"])
|
||||
async def test_get_current_trader_user_accepts_trader_and_admin(role):
|
||||
user = fake_user(role=role)
|
||||
result = await deps.get_current_trader_user(current_user=user)
|
||||
assert result is user
|
||||
|
||||
|
||||
async def test_get_current_trader_user_rejects_viewer():
|
||||
with pytest.raises(ForbiddenException):
|
||||
await deps.get_current_trader_user(current_user=fake_user(role="viewer"))
|
||||
|
||||
|
||||
async def test_get_current_viewer_user_accepts_any_active_role():
|
||||
user = fake_user(role="viewer")
|
||||
result = await deps.get_current_viewer_user(current_user=user)
|
||||
assert result is user
|
||||
|
||||
|
||||
async def test_get_current_admin_user_accepts_admin_role_or_flag():
|
||||
result = await deps.get_current_admin_user(current_user=fake_user(role="admin", is_admin=False))
|
||||
assert result is not None
|
||||
result = await deps.get_current_admin_user(current_user=fake_user(role="trader", is_admin=True))
|
||||
assert result is not None
|
||||
|
||||
|
||||
async def test_get_current_admin_user_rejects_non_admin():
|
||||
with pytest.raises(ForbiddenException):
|
||||
await deps.get_current_admin_user(current_user=fake_user(role="viewer", is_admin=False))
|
||||
|
||||
|
||||
def test_place_order_requires_trader_role_not_bare_auth():
|
||||
"""Regression guard for fix (b).
|
||||
|
||||
Before the fix, POST /orders/place only depended on get_current_user
|
||||
(any authenticated user, including 'viewer'). Assert the route now wires
|
||||
up get_current_trader_user so a future refactor can't silently reintroduce
|
||||
the gap.
|
||||
"""
|
||||
from app.api.v1 import orders
|
||||
|
||||
sig = inspect.signature(orders.place_order)
|
||||
current_user_default = sig.parameters["current_user"].default
|
||||
assert current_user_default.dependency is deps.get_current_trader_user, (
|
||||
"place_order must depend on get_current_trader_user, "
|
||||
"not a weaker dependency — viewers must not be able to place real orders"
|
||||
)
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Tests for app/services/risk_manager.py — Kelly position sizing and
|
||||
regime-adaptive stop-loss/take-profit. Pure computation, no DB/IO needed,
|
||||
but this is the code that directly decides how much real money is risked
|
||||
per trade, so it deserves solid coverage.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services.risk_manager import (
|
||||
REGIME_MULTIPLIERS,
|
||||
AdaptiveSLTPOptimizer,
|
||||
DynamicKellySizer,
|
||||
)
|
||||
|
||||
|
||||
class TestDynamicKellySizerComputeKellyPct:
|
||||
def setup_method(self):
|
||||
self.sizer = DynamicKellySizer(kelly_fraction=0.25)
|
||||
|
||||
def test_zero_when_no_edge_win_rate_too_low(self):
|
||||
# p*b - q < 0 when win rate is low relative to payoff ratio
|
||||
pct = self.sizer.compute_kelly_pct(win_rate=0.1, avg_win=1.0, avg_loss=2.0, confidence=1.0)
|
||||
assert pct == 0.0
|
||||
|
||||
def test_zero_when_avg_loss_is_zero(self):
|
||||
assert self.sizer.compute_kelly_pct(win_rate=0.6, avg_win=3.0, avg_loss=0.0) == 0.0
|
||||
|
||||
def test_zero_when_win_rate_is_zero(self):
|
||||
assert self.sizer.compute_kelly_pct(win_rate=0.0, avg_win=3.0, avg_loss=2.0) == 0.0
|
||||
|
||||
def test_positive_edge_gives_positive_fraction(self):
|
||||
# p=0.6, b=1.5 (avg_win/avg_loss) -> kelly_f = (0.6*1.5 - 0.4)/1.5 = 0.333...
|
||||
pct = self.sizer.compute_kelly_pct(win_rate=0.6, avg_win=3.0, avg_loss=2.0, confidence=1.0)
|
||||
expected_full_kelly = (0.6 * 1.5 - 0.4) / 1.5
|
||||
assert pct == pytest.approx(expected_full_kelly * 0.25, rel=1e-6)
|
||||
|
||||
def test_result_is_clamped_to_50_percent_before_fraction(self):
|
||||
# Even with an absurdly favorable edge, full-Kelly is clamped to 0.5
|
||||
# before applying the 0.25 fractional multiplier -> max output 0.125.
|
||||
pct = self.sizer.compute_kelly_pct(win_rate=0.99, avg_win=100.0, avg_loss=0.1, confidence=1.0)
|
||||
assert pct <= 0.5 * 0.25 + 1e-9
|
||||
|
||||
def test_confidence_scales_output_linearly(self):
|
||||
full_conf = self.sizer.compute_kelly_pct(win_rate=0.6, avg_win=3.0, avg_loss=2.0, confidence=1.0)
|
||||
half_conf = self.sizer.compute_kelly_pct(win_rate=0.6, avg_win=3.0, avg_loss=2.0, confidence=0.5)
|
||||
assert half_conf == pytest.approx(full_conf * 0.5, rel=1e-9)
|
||||
|
||||
def test_output_never_negative(self):
|
||||
pct = self.sizer.compute_kelly_pct(win_rate=0.2, avg_win=1.0, avg_loss=5.0, confidence=1.0)
|
||||
assert pct >= 0.0
|
||||
|
||||
|
||||
class TestDynamicKellySizerVolatilityAdjustedSize:
|
||||
def setup_method(self):
|
||||
self.sizer = DynamicKellySizer()
|
||||
|
||||
def test_floors_at_one_dollar(self):
|
||||
size = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1"), atr_pct=Decimal("20"), max_risk_pct=Decimal("0.01"), regime="choppy"
|
||||
)
|
||||
assert size >= Decimal("1")
|
||||
|
||||
def test_higher_volatility_reduces_size_for_same_regime(self):
|
||||
low_vol = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("1"), regime="neutral"
|
||||
)
|
||||
high_vol = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("10"), regime="neutral"
|
||||
)
|
||||
assert high_vol < low_vol
|
||||
|
||||
def test_trending_regime_sizes_larger_than_choppy(self):
|
||||
trending = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("3"), regime="trending"
|
||||
)
|
||||
choppy = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("3"), regime="choppy"
|
||||
)
|
||||
assert trending > choppy
|
||||
|
||||
def test_unknown_regime_falls_back_to_neutral_factor(self):
|
||||
unknown = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("3"), regime="does-not-exist"
|
||||
)
|
||||
neutral = self.sizer.compute_volatility_adjusted_size(
|
||||
base_size=Decimal("1000"), atr_pct=Decimal("3"), regime="neutral"
|
||||
)
|
||||
assert unknown == neutral
|
||||
|
||||
|
||||
class TestAdaptiveSLTPOptimizerComputeSlTp:
|
||||
def setup_method(self):
|
||||
self.opt = AdaptiveSLTPOptimizer()
|
||||
|
||||
def test_long_stop_loss_is_below_entry_and_tp_above(self):
|
||||
result = self.opt.compute_sl_tp(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
|
||||
assert result["stop_loss"] < 50_000.0 < result["take_profit"]
|
||||
|
||||
def test_short_stop_loss_is_above_entry_and_tp_below(self):
|
||||
result = self.opt.compute_sl_tp(atr=100.0, entry_price=50_000.0, regime="trending", direction="SHORT")
|
||||
assert result["take_profit"] < 50_000.0 < result["stop_loss"]
|
||||
|
||||
def test_direction_is_case_insensitive(self):
|
||||
upper = self.opt.compute_sl_tp(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
|
||||
lower = self.opt.compute_sl_tp(atr=100.0, entry_price=50_000.0, regime="trending", direction="long")
|
||||
assert upper == lower
|
||||
|
||||
@pytest.mark.parametrize("regime", list(REGIME_MULTIPLIERS.keys()))
|
||||
def test_uses_documented_multipliers_for_every_regime(self, regime):
|
||||
result = self.opt.compute_sl_tp(atr=50.0, entry_price=1000.0, regime=regime, direction="LONG")
|
||||
params = REGIME_MULTIPLIERS[regime]
|
||||
assert result["sl_multiplier"] == params["sl"]
|
||||
assert result["tp_multiplier"] == params["tp"]
|
||||
|
||||
def test_choppy_regime_is_never_acceptable(self):
|
||||
"""'choppy' has tp=0.0 and min_rr=99.0 by design — it should never pass
|
||||
the risk/reward gate, which is how the system encodes "don't trade
|
||||
choppy markets" per ARCHITECTURE.md."""
|
||||
result = self.opt.compute_sl_tp(atr=50.0, entry_price=1000.0, regime="choppy", direction="LONG")
|
||||
assert result["acceptable"] is False
|
||||
|
||||
def test_trending_regime_with_good_rr_is_acceptable(self):
|
||||
result = self.opt.compute_sl_tp(atr=100.0, entry_price=50_000.0, regime="trending", direction="LONG")
|
||||
assert result["risk_reward"] >= REGIME_MULTIPLIERS["trending"]["min_rr"]
|
||||
assert result["acceptable"] is True
|
||||
|
||||
def test_unknown_regime_falls_back_to_neutral(self):
|
||||
unknown = self.opt.compute_sl_tp(atr=50.0, entry_price=1000.0, regime="bogus", direction="LONG")
|
||||
neutral = self.opt.compute_sl_tp(atr=50.0, entry_price=1000.0, regime="neutral", direction="LONG")
|
||||
assert unknown == neutral
|
||||
|
||||
|
||||
class TestAdaptiveSLTPOptimizerPartialTpLevels:
|
||||
def setup_method(self):
|
||||
self.opt = AdaptiveSLTPOptimizer()
|
||||
|
||||
def test_returns_two_levels_summing_close_percentage_below_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
|
||||
total_close_pct = sum(lvl["close_percentage"] for lvl in levels)
|
||||
assert 0 < total_close_pct <= 1.0
|
||||
|
||||
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)
|
||||
|
||||
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)
|
||||
@@ -0,0 +1,54 @@
|
||||
"""Tests for fix (d): API key encryption migrated from AES-256-CBC to
|
||||
AES-256-GCM (authenticated encryption), while staying able to decrypt
|
||||
secrets that were already encrypted with the old CBC scheme.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
import pytest
|
||||
from cryptography.hazmat.backends import default_backend
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
|
||||
from app.core.security import decrypt_api_key, encrypt_api_key
|
||||
|
||||
TEST_KEY = "11" * 32 # 32-byte hex key
|
||||
|
||||
|
||||
def _legacy_cbc_encrypt(plaintext: str, key_hex: str) -> tuple[str, str]:
|
||||
"""Recreate the old (pre-fix) AES-256-CBC scheme to seed a 'legacy' ciphertext."""
|
||||
key = bytes.fromhex(key_hex)
|
||||
iv = uuid.uuid4().bytes # 16 bytes
|
||||
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=default_backend())
|
||||
encryptor = cipher.encryptor()
|
||||
plaintext_bytes = plaintext.encode("utf-8")
|
||||
pad_len = 16 - (len(plaintext_bytes) % 16)
|
||||
padded = plaintext_bytes + bytes([pad_len] * pad_len)
|
||||
ciphertext = encryptor.update(padded) + encryptor.finalize()
|
||||
return ciphertext.hex(), iv.hex()
|
||||
|
||||
|
||||
def test_encrypt_then_decrypt_roundtrip():
|
||||
ciphertext_hex, nonce_hex = encrypt_api_key("super-secret-key", key_hex=TEST_KEY)
|
||||
assert decrypt_api_key(ciphertext_hex, nonce_hex, key_hex=TEST_KEY) == "super-secret-key"
|
||||
|
||||
|
||||
def test_new_encryption_uses_gcm_12_byte_nonce():
|
||||
_, nonce_hex = encrypt_api_key("anything", key_hex=TEST_KEY)
|
||||
assert len(bytes.fromhex(nonce_hex)) == 12, "new secrets must use a 12-byte GCM nonce, not the old 16-byte CBC IV"
|
||||
|
||||
|
||||
def test_tampered_gcm_ciphertext_is_rejected():
|
||||
"""The whole point of GCM over CBC: a flipped bit must be detected, not silently decrypt to garbage."""
|
||||
ciphertext_hex, nonce_hex = encrypt_api_key("super-secret-key", key_hex=TEST_KEY)
|
||||
tampered = bytearray(bytes.fromhex(ciphertext_hex))
|
||||
tampered[0] ^= 0xFF
|
||||
with pytest.raises(Exception):
|
||||
decrypt_api_key(tampered.hex(), nonce_hex, key_hex=TEST_KEY)
|
||||
|
||||
|
||||
def test_legacy_cbc_ciphertexts_still_decrypt():
|
||||
"""Credentials encrypted before the GCM migration must not be bricked."""
|
||||
ciphertext_hex, iv_hex = _legacy_cbc_encrypt("old-style-secret", TEST_KEY)
|
||||
assert len(bytes.fromhex(iv_hex)) == 16
|
||||
assert decrypt_api_key(ciphertext_hex, iv_hex, key_hex=TEST_KEY) == "old-style-secret"
|
||||
@@ -0,0 +1,208 @@
|
||||
"""Tests for the pure scoring/classification helpers in
|
||||
app/services/signal_service.py — the 13-algorithm voting core that decides
|
||||
BUY/SELL/STRONG signals. These functions take plain indicator dicts/lists
|
||||
and return classifications; no DB or network I/O involved.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from decimal import Decimal
|
||||
|
||||
from app.services.signal_service import (
|
||||
BUY,
|
||||
CAUTION_LONG,
|
||||
CAUTION_SHORT,
|
||||
SELL,
|
||||
SQUEEZE_ALERT,
|
||||
STRONG_BUY,
|
||||
STRONG_SELL,
|
||||
_calculate_pnl,
|
||||
_classify_signal_bb,
|
||||
_classify_signal_combined,
|
||||
_detect_squeeze,
|
||||
_get_bb_values,
|
||||
_get_rsi_values,
|
||||
_get_sma_values,
|
||||
)
|
||||
|
||||
|
||||
class TestExtractors:
|
||||
def test_get_bb_values_returns_none_when_missing(self):
|
||||
assert _get_bb_values({}) is None
|
||||
|
||||
def test_get_bb_values_returns_none_when_incomplete(self):
|
||||
assert _get_bb_values({"bollinger_bands": {"upper": [1]}}) is None
|
||||
|
||||
def test_get_bb_values_returns_dict_when_complete(self):
|
||||
bb = {"upper": [1], "middle": [0], "lower": [-1]}
|
||||
assert _get_bb_values({"bollinger_bands": bb}) == bb
|
||||
|
||||
def test_get_rsi_values_returns_none_for_empty_list(self):
|
||||
assert _get_rsi_values({"rsi_14": []}) is None
|
||||
|
||||
def test_get_rsi_values_returns_list(self):
|
||||
assert _get_rsi_values({"rsi_14": [55.0]}) == [55.0]
|
||||
|
||||
def test_get_sma_values_returns_none_when_missing(self):
|
||||
assert _get_sma_values({}) is None
|
||||
|
||||
def test_get_sma_values_returns_list(self):
|
||||
assert _get_sma_values({"sma_20": [100.0]}) == [100.0]
|
||||
|
||||
|
||||
class TestDetectSqueeze:
|
||||
def test_false_when_not_enough_history(self):
|
||||
bb = {"upper": [110] * 3, "lower": [90] * 3}
|
||||
assert _detect_squeeze(bb, lookback=10) is False
|
||||
|
||||
def test_true_when_inner_bands_are_narrow_relative_to_outer(self):
|
||||
bb = {
|
||||
"upper": [110] * 10,
|
||||
"lower": [90] * 10,
|
||||
"upper_1": [100.5],
|
||||
"lower_1": [99.5],
|
||||
}
|
||||
assert _detect_squeeze(bb) is True
|
||||
|
||||
def test_false_when_bands_are_wide_and_stable(self):
|
||||
bb = {
|
||||
"upper": [120, 119, 118, 117, 116, 115, 114, 113, 112, 111],
|
||||
"lower": [80, 79, 78, 77, 76, 75, 74, 73, 72, 71],
|
||||
}
|
||||
# width shrinks steadily from 40 to 40 (constant) -> current == min -> squeeze True by the
|
||||
# "at the minimum of the lookback" rule. Use a widening series instead to assert False.
|
||||
bb_widening = {
|
||||
"upper": [101, 102, 103, 104, 105, 106, 107, 108, 109, 130],
|
||||
"lower": [99, 98, 97, 96, 95, 94, 93, 92, 91, 70],
|
||||
}
|
||||
assert _detect_squeeze(bb_widening) is False
|
||||
|
||||
|
||||
class TestClassifySignalBb:
|
||||
BASE_BB = {"upper": [110], "lower": [90], "upper_1": [105], "lower_1": [95], "middle": [100]}
|
||||
|
||||
def test_none_when_price_inside_bands(self):
|
||||
signal_type, strength = _classify_signal_bb(100.0, self.BASE_BB, rsi=[50], sma=[100])
|
||||
assert (signal_type, strength) == (None, None)
|
||||
|
||||
def test_strong_buy_when_price_above_upper2_rsi_bullish_sma_above_middle(self):
|
||||
signal_type, strength = _classify_signal_bb(115.0, self.BASE_BB, rsi=[65], sma=[102])
|
||||
assert (signal_type, strength) == (STRONG_BUY, "STRONG")
|
||||
|
||||
def test_caution_short_when_price_above_upper2_and_rsi_overbought(self):
|
||||
signal_type, strength = _classify_signal_bb(115.0, self.BASE_BB, rsi=[80], sma=[102])
|
||||
assert (signal_type, strength) == (CAUTION_SHORT, "MODERATE")
|
||||
|
||||
def test_moderate_buy_when_price_above_upper1_but_below_upper2(self):
|
||||
signal_type, strength = _classify_signal_bb(107.0, self.BASE_BB, rsi=[60], sma=[100])
|
||||
assert (signal_type, strength) == (BUY, "MODERATE")
|
||||
|
||||
def test_strong_sell_when_price_below_lower2_rsi_bearish_sma_below_middle(self):
|
||||
signal_type, strength = _classify_signal_bb(85.0, self.BASE_BB, rsi=[35], sma=[98])
|
||||
assert (signal_type, strength) == (STRONG_SELL, "STRONG")
|
||||
|
||||
def test_caution_long_when_price_below_lower2_and_rsi_oversold(self):
|
||||
signal_type, strength = _classify_signal_bb(85.0, self.BASE_BB, rsi=[20], sma=[98])
|
||||
assert (signal_type, strength) == (CAUTION_LONG, "MODERATE")
|
||||
|
||||
|
||||
class TestClassifySignalCombined:
|
||||
NO_SQUEEZE_BB = {"upper": [110], "lower": [90], "upper_1": [105], "lower_1": [95], "middle": [100]}
|
||||
|
||||
def test_invalid_close_price_returns_neutral(self):
|
||||
result = _classify_signal_combined(
|
||||
math.nan, self.NO_SQUEEZE_BB, None, None, None, None, None
|
||||
)
|
||||
assert result == (None, None, 0.0, {})
|
||||
|
||||
result_neg = _classify_signal_combined(
|
||||
-5.0, self.NO_SQUEEZE_BB, None, None, None, None, None
|
||||
)
|
||||
assert result_neg == (None, None, 0.0, {})
|
||||
|
||||
def test_squeeze_overrides_everything_else(self):
|
||||
squeezed_bb = {
|
||||
"upper": [110] * 10,
|
||||
"lower": [90] * 10,
|
||||
"upper_1": [100.5],
|
||||
"lower_1": [99.5],
|
||||
}
|
||||
result = _classify_signal_combined(
|
||||
100.0, squeezed_bb, rsi=[65], sma=[102],
|
||||
macd_data={"macd_line": [1, 2], "signal_line": [0, 0.5]},
|
||||
st_data=None, vol_data=None,
|
||||
)
|
||||
assert result == (SQUEEZE_ALERT, "MODERATE", 0.5, {})
|
||||
|
||||
def test_caution_signal_passes_through_with_fixed_confidence(self):
|
||||
signal_type, strength, confidence, raw_scores = _classify_signal_combined(
|
||||
115.0, self.NO_SQUEEZE_BB, rsi=[80], sma=[102],
|
||||
macd_data=None, st_data=None, vol_data=None,
|
||||
)
|
||||
assert (signal_type, strength, confidence) == (CAUTION_SHORT, "MODERATE", 0.5)
|
||||
assert raw_scores == {}
|
||||
|
||||
def test_all_neutral_inputs_yield_no_signal_and_zero_confidence(self):
|
||||
signal_type, strength, confidence, raw_scores = _classify_signal_combined(
|
||||
100.0, self.NO_SQUEEZE_BB, rsi=[50], sma=[100],
|
||||
macd_data=None, st_data=None, vol_data=None,
|
||||
)
|
||||
assert (signal_type, strength) == (None, None)
|
||||
assert confidence == 0.0
|
||||
assert all(v == 0.0 for v in raw_scores.values())
|
||||
|
||||
def test_single_strong_bb_vote_alone_only_reaches_moderate_buy(self):
|
||||
"""A lone (even 'STRONG') indicator can't clear the combined-score
|
||||
STRONG threshold by itself — the system is designed to require
|
||||
multiple confirming signals for a STRONG classification."""
|
||||
signal_type, strength, _confidence, raw_scores = _classify_signal_combined(
|
||||
115.0, self.NO_SQUEEZE_BB, rsi=[65], sma=[102],
|
||||
macd_data=None, st_data=None, vol_data=None,
|
||||
)
|
||||
assert signal_type == BUY
|
||||
assert strength == "MODERATE"
|
||||
assert raw_scores["double_bb_rsi"] == 2.0
|
||||
|
||||
def test_bullish_alignment_across_groups_never_flips_to_bearish(self):
|
||||
signal_type, _strength, _confidence, raw_scores = _classify_signal_combined(
|
||||
115.0, self.NO_SQUEEZE_BB, rsi=[65], sma=[102],
|
||||
macd_data={"macd_line": [-1, 1], "signal_line": [0, 0]},
|
||||
st_data={"trend": [True]},
|
||||
vol_data=[False, True],
|
||||
)
|
||||
assert signal_type in (BUY, STRONG_BUY)
|
||||
assert raw_scores["double_bb_rsi"] > 0
|
||||
assert raw_scores["macd_crossover"] > 0
|
||||
assert raw_scores["supertrend"] > 0
|
||||
|
||||
def test_bearish_alignment_is_the_mirror_of_bullish(self):
|
||||
signal_type, _strength, _confidence, raw_scores = _classify_signal_combined(
|
||||
85.0, self.NO_SQUEEZE_BB, rsi=[35], sma=[98],
|
||||
macd_data={"macd_line": [1, -1], "signal_line": [0, 0]},
|
||||
st_data={"trend": [False]},
|
||||
vol_data=[False, True],
|
||||
)
|
||||
assert signal_type in (SELL, STRONG_SELL)
|
||||
assert raw_scores["double_bb_rsi"] < 0
|
||||
assert raw_scores["macd_crossover"] < 0
|
||||
assert raw_scores["supertrend"] < 0
|
||||
|
||||
def test_disabled_strategies_are_zeroed_out(self):
|
||||
_signal_type, _strength, _confidence, raw_scores = _classify_signal_combined(
|
||||
115.0, self.NO_SQUEEZE_BB, rsi=[65], sma=[102],
|
||||
macd_data=None, st_data=None, vol_data=None,
|
||||
enabled_strategies=["macd_crossover"], # double_bb_rsi not in the allow-list
|
||||
)
|
||||
assert raw_scores["double_bb_rsi"] == 0.0
|
||||
|
||||
|
||||
class TestCalculatePnl:
|
||||
def test_long_profit(self):
|
||||
pnl, pnl_pct = _calculate_pnl(Decimal("100"), Decimal("110"), "LONG", Decimal("2"))
|
||||
assert pnl == Decimal("20")
|
||||
assert pnl_pct == Decimal("10")
|
||||
|
||||
def test_short_profit(self):
|
||||
pnl, pnl_pct = _calculate_pnl(Decimal("100"), Decimal("90"), "SHORT", Decimal("2"))
|
||||
assert pnl == Decimal("20")
|
||||
assert pnl_pct == Decimal("10")
|
||||
@@ -0,0 +1,234 @@
|
||||
"""Tests for app/services/trade_executor.py — this is the code that actually
|
||||
opens/closes (paper) trades from signals, so its rules matter a lot:
|
||||
1. Only STRONG_BUY/STRONG_SELL open new trades (signal != trade).
|
||||
2. A same-direction open trade blocks a duplicate entry (dedup).
|
||||
3. An opposite-direction open trade gets closed on a STRONG reversal.
|
||||
4. Extreme ATR% (>8% or <0.5%) skips the entry (volatility filter).
|
||||
5. Hitting MAX_OPEN_TRADES evicts the oldest trade when all are winners
|
||||
(hybrid eviction — FIFO fallback).
|
||||
|
||||
Runs against a real (in-memory SQLite) DB, exercising the actual SQL
|
||||
filters `execute_signal_trade` builds, not a mocked stand-in for them.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models import HypotheticalTrade, Signal, User
|
||||
from app.services.trade_executor import (
|
||||
MAX_OPEN_TRADES,
|
||||
STRONG_BUY,
|
||||
STRONG_SELL,
|
||||
_calculate_pnl,
|
||||
execute_signal_trade,
|
||||
)
|
||||
|
||||
|
||||
def make_user(**prefs_overrides) -> User:
|
||||
prefs = {"trade_size": 10, "auto_trade_tokens": []}
|
||||
prefs.update(prefs_overrides)
|
||||
return User(
|
||||
id=uuid.uuid4(),
|
||||
username=f"user-{uuid.uuid4().hex[:8]}",
|
||||
email=f"{uuid.uuid4().hex[:8]}@example.com",
|
||||
password_hash="x",
|
||||
role="trader",
|
||||
is_active=True,
|
||||
preferences=prefs,
|
||||
)
|
||||
|
||||
|
||||
def make_signal(signal_type: str, symbol: str = "BTC/USDT", atr_pct_snapshot=None) -> Signal:
|
||||
snapshot = None
|
||||
if atr_pct_snapshot is not None:
|
||||
snapshot = json.dumps({"atr_14": [atr_pct_snapshot["atr_abs"]], "confidence": 0.5})
|
||||
return Signal(
|
||||
symbol=symbol,
|
||||
exchange="mexc",
|
||||
timeframe="1h",
|
||||
signal_type=signal_type,
|
||||
strength="STRONG",
|
||||
price=Decimal("100"),
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
indicators_snapshot=snapshot,
|
||||
)
|
||||
|
||||
|
||||
async def _open_trades_for(db_session, user_id, symbol=None):
|
||||
query = select(HypotheticalTrade).where(HypotheticalTrade.user_id == user_id)
|
||||
if symbol:
|
||||
query = query.where(HypotheticalTrade.symbol == symbol)
|
||||
result = await db_session.execute(query)
|
||||
return result.scalars().all()
|
||||
|
||||
|
||||
class TestCalculatePnl:
|
||||
def test_long_profit_when_price_rises(self):
|
||||
pnl, pnl_pct = _calculate_pnl(Decimal("100"), Decimal("110"), "LONG", Decimal("2"))
|
||||
assert pnl == Decimal("20")
|
||||
assert pnl_pct == Decimal("10")
|
||||
|
||||
def test_short_profit_when_price_falls(self):
|
||||
pnl, pnl_pct = _calculate_pnl(Decimal("100"), Decimal("90"), "SHORT", Decimal("2"))
|
||||
assert pnl == Decimal("20")
|
||||
assert pnl_pct == Decimal("10")
|
||||
|
||||
def test_long_loss_when_price_falls(self):
|
||||
pnl, _ = _calculate_pnl(Decimal("100"), Decimal("90"), "LONG", Decimal("1"))
|
||||
assert pnl == Decimal("-10")
|
||||
|
||||
|
||||
async def test_non_strong_signal_does_nothing(db_session):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
|
||||
signal = make_signal("BUY") # not STRONG_BUY
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("100"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id)
|
||||
assert trades == []
|
||||
|
||||
|
||||
async def test_strong_buy_opens_long_trade_for_eligible_user(db_session):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
signal = make_signal(STRONG_BUY)
|
||||
db_session.add(signal)
|
||||
await db_session.flush()
|
||||
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("50000"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id, "BTC/USDT")
|
||||
assert len(trades) == 1
|
||||
assert trades[0].direction == "LONG"
|
||||
assert trades[0].status == "OPEN"
|
||||
assert trades[0].entry_price == Decimal("50000")
|
||||
|
||||
|
||||
async def test_user_not_subscribed_to_token_is_skipped(db_session):
|
||||
user = make_user(auto_trade_tokens=["ETH/USDT"]) # only wants ETH, signal is for BTC
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
signal = make_signal(STRONG_BUY, symbol="BTC/USDT")
|
||||
db_session.add(signal)
|
||||
await db_session.flush()
|
||||
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("50000"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id)
|
||||
assert trades == []
|
||||
|
||||
|
||||
async def test_dedup_skips_when_same_direction_already_open(db_session):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
|
||||
existing = HypotheticalTrade(
|
||||
user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h",
|
||||
direction="LONG", entry_price=Decimal("40000"),
|
||||
entry_time=datetime.now(timezone.utc), quantity=Decimal("1"), status="OPEN",
|
||||
)
|
||||
db_session.add(existing)
|
||||
await db_session.flush()
|
||||
|
||||
signal = make_signal(STRONG_BUY, symbol="BTC/USDT")
|
||||
db_session.add(signal)
|
||||
await db_session.flush()
|
||||
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("50000"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id, "BTC/USDT")
|
||||
assert len(trades) == 1, "must not open a second trade in the same direction"
|
||||
assert trades[0].entry_price == Decimal("40000"), "the original open trade must be untouched"
|
||||
assert trades[0].status == "OPEN"
|
||||
|
||||
|
||||
async def test_strong_reversal_closes_opposite_trade_and_opens_new_one(db_session):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
|
||||
existing_short = HypotheticalTrade(
|
||||
user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h",
|
||||
direction="SHORT", entry_price=Decimal("60000"),
|
||||
entry_time=datetime.now(timezone.utc), quantity=Decimal("1"), status="OPEN",
|
||||
)
|
||||
db_session.add(existing_short)
|
||||
await db_session.flush()
|
||||
|
||||
signal = make_signal(STRONG_BUY, symbol="BTC/USDT") # LONG signal reverses the SHORT
|
||||
db_session.add(signal)
|
||||
await db_session.flush()
|
||||
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("50000"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id, "BTC/USDT")
|
||||
assert len(trades) == 2
|
||||
closed = [t for t in trades if t.status == "CLOSED"]
|
||||
opened = [t for t in trades if t.status == "OPEN"]
|
||||
assert len(closed) == 1 and closed[0].exit_reason == "REVERSAL"
|
||||
assert closed[0].direction == "SHORT"
|
||||
assert len(opened) == 1 and opened[0].direction == "LONG"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("atr_abs,expected_reason", [(10.0, "too high"), (0.1, "too low")])
|
||||
async def test_volatility_filter_skips_extreme_atr(db_session, atr_abs, expected_reason):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
|
||||
# current_price=100 -> atr_pct = atr_abs/100*100 = atr_abs (%)
|
||||
signal = make_signal(STRONG_BUY, symbol="BTC/USDT", atr_pct_snapshot={"atr_abs": atr_abs})
|
||||
db_session.add(signal)
|
||||
await db_session.flush()
|
||||
|
||||
await execute_signal_trade(db_session, signal, "BTC/USDT", "mexc", "1h", Decimal("100"))
|
||||
|
||||
trades = await _open_trades_for(db_session, user.id)
|
||||
assert trades == [], f"must skip entry when ATR% is {expected_reason} ({atr_abs}%)"
|
||||
|
||||
|
||||
async def test_hybrid_eviction_evicts_oldest_when_all_open_trades_are_winners(db_session):
|
||||
user = make_user()
|
||||
db_session.add(user)
|
||||
await db_session.flush()
|
||||
|
||||
base_time = datetime.now(timezone.utc) - timedelta(days=1)
|
||||
existing_trades = []
|
||||
for i in range(MAX_OPEN_TRADES):
|
||||
t = HypotheticalTrade(
|
||||
user_id=user.id, symbol=f"SYM{i}/USDT", exchange="mexc", timeframe="1h",
|
||||
direction="LONG", entry_price=Decimal("100"),
|
||||
entry_time=base_time + timedelta(minutes=i), # SYM0 is oldest
|
||||
quantity=Decimal("1"), status="OPEN",
|
||||
)
|
||||
existing_trades.append(t)
|
||||
db_session.add_all(existing_trades)
|
||||
await db_session.flush()
|
||||
|
||||
# Incoming signal for a brand-new symbol; current_price=200 > entry_price=100
|
||||
# for every existing trade -> all of them are winners at this price.
|
||||
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("200"))
|
||||
|
||||
all_trades = await _open_trades_for(db_session, user.id)
|
||||
open_trades = [t for t in all_trades if t.status == "OPEN"]
|
||||
closed_trades = [t for t in all_trades if t.status == "CLOSED"]
|
||||
|
||||
assert len(open_trades) == MAX_OPEN_TRADES, "9 survivors + 1 newly opened = MAX_OPEN_TRADES"
|
||||
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)
|
||||
Reference in New Issue
Block a user