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:
Le
2026-07-03 21:34:18 +07:00
parent b37fa1804e
commit afbe4f4f15
12 changed files with 1046 additions and 0 deletions
+33
View File
@@ -0,0 +1,33 @@
name: Backend tests
on:
push:
branches: [master]
paths:
- "backend/**"
- ".gitea/workflows/backend-tests.yml"
pull_request:
paths:
- "backend/**"
- ".gitea/workflows/backend-tests.yml"
jobs:
pytest:
runs-on: ubuntu-latest
defaults:
run:
working-directory: backend
steps:
- uses: actions/checkout@v4
- uses: actions/setup-python@v5
with:
python-version: "3.13"
- name: Install dependencies
run: |
python -m pip install --upgrade pip
pip install -r requirements.txt -r requirements-dev.txt
- name: Run pytest
run: python -m pytest tests/ -v
+3
View File
@@ -0,0 +1,3 @@
[pytest]
asyncio_mode = auto
testpaths = tests
+5
View File
@@ -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
View File
+75
View File
@@ -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()
+47
View File
@@ -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"
+78
View File
@@ -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"
)
+154
View File
@@ -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)
+54
View File
@@ -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")
+234
View File
@@ -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)