diff --git a/.gitea/workflows/backend-tests.yml b/.gitea/workflows/backend-tests.yml new file mode 100644 index 0000000..09070f1 --- /dev/null +++ b/.gitea/workflows/backend-tests.yml @@ -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 diff --git a/backend/pytest.ini b/backend/pytest.ini new file mode 100644 index 0000000..78c5011 --- /dev/null +++ b/backend/pytest.ini @@ -0,0 +1,3 @@ +[pytest] +asyncio_mode = auto +testpaths = tests diff --git a/backend/requirements-dev.txt b/backend/requirements-dev.txt new file mode 100644 index 0000000..e459454 --- /dev/null +++ b/backend/requirements-dev.txt @@ -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 diff --git a/backend/tests/__init__.py b/backend/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py new file mode 100644 index 0000000..10964a3 --- /dev/null +++ b/backend/tests/conftest.py @@ -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() diff --git a/backend/tests/test_cors_config.py b/backend/tests/test_cors_config.py new file mode 100644 index 0000000..985ebeb --- /dev/null +++ b/backend/tests/test_cors_config.py @@ -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) diff --git a/backend/tests/test_order_exchange_routing.py b/backend/tests/test_order_exchange_routing.py new file mode 100644 index 0000000..e625a78 --- /dev/null +++ b/backend/tests/test_order_exchange_routing.py @@ -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" diff --git a/backend/tests/test_rbac_deps.py b/backend/tests/test_rbac_deps.py new file mode 100644 index 0000000..aab9922 --- /dev/null +++ b/backend/tests/test_rbac_deps.py @@ -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" + ) diff --git a/backend/tests/test_risk_manager.py b/backend/tests/test_risk_manager.py new file mode 100644 index 0000000..62c5666 --- /dev/null +++ b/backend/tests/test_risk_manager.py @@ -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) diff --git a/backend/tests/test_security_encryption.py b/backend/tests/test_security_encryption.py new file mode 100644 index 0000000..81c9a22 --- /dev/null +++ b/backend/tests/test_security_encryption.py @@ -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" diff --git a/backend/tests/test_signal_service_scoring.py b/backend/tests/test_signal_service_scoring.py new file mode 100644 index 0000000..a295e4e --- /dev/null +++ b/backend/tests/test_signal_service_scoring.py @@ -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") diff --git a/backend/tests/test_trade_executor.py b/backend/tests/test_trade_executor.py new file mode 100644 index 0000000..70a1ac5 --- /dev/null +++ b/backend/tests/test_trade_executor.py @@ -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)