diff --git a/backend/tests/test_indicator_service.py b/backend/tests/test_indicator_service.py new file mode 100644 index 0000000..f24e86f --- /dev/null +++ b/backend/tests/test_indicator_service.py @@ -0,0 +1,277 @@ +"""Tests for app/services/indicator_service.py — the pure-math indicator +library that every algorithm in the 13-vote signal system is built on top +of. Runs against the pure-Python fallback path (this venv has no numpy +installed), which is the code path actually exercised in a minimal +deployment of this backend. +""" +from __future__ import annotations + +import math + +import pytest + +from app.services.indicator_service import ( + atr, + bollinger_bands, + detect_market_regime, + ema, + macd, + mfi, + obv, + obv_signal, + rsi, + sma, + vwap, +) + + +def candle(high, low, close, volume=100.0): + return {"high": high, "low": low, "close": close, "volume": volume} + + +class TestSma: + def test_leading_entries_are_none(self): + result = sma([1, 2, 3, 4, 5], period=3) + assert result[:2] == [None, None] + + def test_values_match_hand_computed_average(self): + result = sma([1, 2, 3, 4, 5], period=3) + assert result == [None, None, 2.0, 3.0, 4.0] + + def test_rejects_non_positive_period(self): + with pytest.raises(ValueError): + sma([1, 2, 3], period=0) + + +class TestEma: + def test_seeded_with_sma_then_smoothed(self): + # period=3 -> multiplier=0.5; seed = mean([1,2,3])=2.0 + result = ema([1, 2, 3, 4, 5], period=3) + assert result[2] == pytest.approx(2.0) + assert result[3] == pytest.approx((4 - 2.0) * 0.5 + 2.0) # 3.0 + assert result[4] == pytest.approx((5 - 3.0) * 0.5 + 3.0) # 4.0 + + +class TestRsi: + def test_all_gains_yields_high_rsi_using_sentinel_rs(self): + prices = list(range(1, 17)) # strictly increasing, 15 deltas, all +1 + result = rsi(prices, period=14) + assert all(v is None for v in result[:14]) + # avg_loss stays 0 -> code uses RS=100 sentinel -> RSI = 100 - 100/101 + expected = 100.0 - 100.0 / 101.0 + assert result[14] == pytest.approx(expected) + + def test_all_losses_yields_rsi_near_zero(self): + prices = list(range(16, 0, -1)) # strictly decreasing + result = rsi(prices, period=14) + assert result[14] == pytest.approx(0.0) + + def test_flat_prices_do_not_yield_neutral_50(self): + """Documents a discovered quirk (not fixed here — flagged for the + team to decide on): when there's truly zero price movement, the + code sets an internal `rs = 50.0` sentinel intending "neutral", but + that value is still run through the RSI formula + (100 - 100/(1+rs)), which maps rs=50 to RSI≈98.04, not the + conventionally-expected neutral RSI of 50. A perfectly flat run + (e.g. an illiquid pair or stablecoin) would misreport as + near-overbought instead of neutral.""" + prices = [10.0] * 16 + result = rsi(prices, period=14) + assert result[14] == pytest.approx(100.0 - 100.0 / 51.0) + + def test_insufficient_data_returns_all_none(self): + result = rsi([1, 2, 3], period=14) + assert result == [None, None, None] + + +class TestBollingerBands: + def test_zero_variance_collapses_bands_to_middle(self): + bands = bollinger_bands([10.0] * 5, period=5, std_dev=2.0) + assert bands["middle"][-1] == pytest.approx(10.0) + assert bands["upper"][-1] == pytest.approx(10.0) + assert bands["lower"][-1] == pytest.approx(10.0) + + def test_hand_computed_bands_for_linear_prices(self): + bands = bollinger_bands([1.0, 2.0, 3.0, 4.0, 5.0], period=5, std_dev=2.0) + sd = math.sqrt(2.0) # population variance of [1..5] around mean 3 = 2.0 + assert bands["middle"][-1] == pytest.approx(3.0) + assert bands["upper"][-1] == pytest.approx(3.0 + 2 * sd) + assert bands["lower"][-1] == pytest.approx(3.0 - 2 * sd) + assert bands["upper_1"][-1] == pytest.approx(3.0 + sd) + assert bands["lower_1"][-1] == pytest.approx(3.0 - sd) + + def test_upper_always_above_lower_when_present(self): + bands = bollinger_bands([5, 3, 8, 1, 9, 2, 7, 4, 6, 10.0], period=5) + for u, l in zip(bands["upper"], bands["lower"]): + if u is not None and l is not None: + assert u >= l + + +class TestMacd: + def test_rejects_fast_not_less_than_slow(self): + with pytest.raises(ValueError): + macd([1.0] * 30, fast=26, slow=12) + + def test_returns_expected_keys_with_matching_length(self): + prices = [float(i % 7 + 10) for i in range(40)] + result = macd(prices, fast=5, slow=10, signal=3) + assert set(result.keys()) == {"macd_line", "signal_line", "histogram"} + assert len(result["macd_line"]) == len(prices) + assert len(result["signal_line"]) == len(prices) + assert len(result["histogram"]) == len(prices) + + def test_macd_line_is_ema_fast_minus_ema_slow(self): + prices = [float(i) for i in range(1, 41)] + result = macd(prices, fast=5, slow=10, signal=3) + ema_fast = ema(prices, 5) + ema_slow = ema(prices, 10) + for i in range(9, len(prices)): + assert result["macd_line"][i] == pytest.approx(ema_fast[i] - ema_slow[i]) + + +class TestAtr: + def test_constant_true_range_gives_exact_atr(self): + # high=105, low=95 always -> TR=10 for every bar (including vs prev close=100) + candles = [candle(105, 95, 100) for _ in range(4)] + result = atr(candles, period=3) + assert result[:2] == [None, None] + assert result[2] == pytest.approx(10.0) + assert result[3] == pytest.approx(10.0) + + def test_insufficient_candles_returns_all_none(self): + assert atr([candle(1, 1, 1)], period=14) == [None] + + +class TestVwap: + def test_cumulative_volume_weighted_average(self): + candles = [candle(10, 8, 9, volume=100), candle(12, 10, 11, volume=50)] + result = vwap(candles) + assert result[0] == pytest.approx(9.0) # typical=(10+8+9)/3=9, only candle so far + assert result[1] == pytest.approx((9.0 * 100 + 11.0 * 50) / 150) + + def test_empty_input_returns_empty_list(self): + assert vwap([]) == [] + + +class TestObv: + def test_hand_computed_cumulative_volume(self): + candles = [ + candle(0, 0, 10, volume=100), + candle(0, 0, 11, volume=200), # up -> +200 + candle(0, 0, 10, volume=150), # down -> -150 + candle(0, 0, 12, volume=300), # up -> +300 + ] + assert obv(candles) == [100.0, 300.0, 150.0, 450.0] + + def test_unchanged_close_keeps_obv_flat(self): + candles = [candle(0, 0, 10, volume=100), candle(0, 0, 10, volume=50)] + assert obv(candles) == [100.0, 100.0] + + +class TestObvSignal: + def test_insufficient_data_returns_all_none(self): + crossovers, sma_vals = obv_signal([1.0, 2.0], period=3) + assert crossovers == [None, None] + assert sma_vals == [None, None] + + def test_detects_bullish_crossover(self): + # OBV sits at/below its SMA (all zero) then jumps above it + crossovers, obv_sma = obv_signal([0, 0, 0, 0, 10], period=3) + assert crossovers[4] is True + + def test_detects_bearish_crossover(self): + crossovers, _ = obv_signal([10, 10, 10, 10, 0], period=3) + assert crossovers[4] is False + + +class TestMfi: + def test_insufficient_candles_returns_all_none(self): + assert mfi([candle(1, 1, 1)], period=14) == [None] + + def test_overbought_when_no_negative_flow(self): + # typical price strictly increasing -> every period contributes only + # positive flow -> neg_flow == 0 -> MFI defined as 100.0. + # Index 4 (not 3) is asserted because index 3 is the very first + # computed value and hits the negative-indexing quirk below. + candles = [candle(10 + i, 10 + i, 10 + i, volume=100) for i in range(5)] + result = mfi(candles, period=3) + assert result[4] == pytest.approx(100.0) + + def test_first_computed_value_has_a_wraparound_indexing_quirk(self): + """Documents a discovered quirk (not fixed here — flagged for the + team to decide on): for the first computed MFI value in a series, + the loop compares `typical_prices[j-1]` with `j=0`, which in Python + wraps around to `typical_prices[-1]` (the LAST candle in the whole + series) instead of having no prior candle to compare against. This + spuriously injects one bogus flow-direction comparison. In practice + this only taints the single oldest computed value in a long series + (never the latest, which is what signal_service.py actually reads), + so real-world impact is negligible — but it is objectively wrong.""" + candles = [candle(10 + i, 10 + i, 10 + i, volume=100) for i in range(5)] + result = mfi(candles, period=3) + # Without the quirk this would also be 100.0 (strictly increasing, + # no real negative flow) — the quirk drags it down to ~69.7. + assert result[3] == pytest.approx(69.69696969696969) + + +class TestDetectMarketRegime: + def _adx(self, value): + return {"adx": [value]} + + def test_squeeze_with_volume_spike_is_breakout(self): + bb = { + "upper": [110.0] * 20, + "lower": [90.0] * 20, + } + # Force squeeze: last width == min width + regime = detect_market_regime( + self._adx(15), bb, atr_pct=1.0, volume_data=[False, True], + ) + assert regime == "breakout" + + def test_squeeze_without_volume_spike_is_squeeze(self): + bb = {"upper": [110.0] * 20, "lower": [90.0] * 20} + regime = detect_market_regime( + self._adx(15), bb, atr_pct=1.0, volume_data=[False, False], + ) + assert regime == "squeeze" + + def test_high_atr_without_squeeze_is_volatile(self): + bb = {"upper": [], "lower": []} + regime = detect_market_regime( + self._adx(15), bb, atr_pct=10.0, volume_data=None, + ) + assert regime == "volatile" + + def test_high_adx_alone_is_not_enough_to_confirm_trending(self): + """High ADX alone yields 'neutral', not 'trending' — the function + also requires Efficiency-Ratio/Choppiness confirmation (which needs + `prices`/`highs`/`lows`), matching its documented "multi-factor" + design rather than being a pure ADX threshold classifier.""" + bb = {"upper": [], "lower": []} + regime = detect_market_regime( + self._adx(30), bb, atr_pct=1.0, volume_data=None, + ) + assert regime == "neutral" + + def test_high_adx_with_confirming_efficiency_ratio_is_trending(self): + bb = {"upper": [], "lower": []} + prices = [float(i) for i in range(1, 30)] # strictly trending -> ER ~1.0 + regime = detect_market_regime( + self._adx(30), bb, atr_pct=1.0, volume_data=None, prices=prices, + ) + assert regime == "trending" + + def test_low_adx_is_sideways(self): + bb = {"upper": [], "lower": []} + regime = detect_market_regime( + self._adx(10), bb, atr_pct=1.0, volume_data=None, + ) + assert regime == "sideways" + + def test_mid_range_with_no_data_is_neutral(self): + bb = {"upper": [], "lower": []} + regime = detect_market_regime( + self._adx(22), bb, atr_pct=1.0, volume_data=None, + ) + assert regime == "neutral" diff --git a/backend/tests/test_signal_service_async.py b/backend/tests/test_signal_service_async.py new file mode 100644 index 0000000..e08f23d --- /dev/null +++ b/backend/tests/test_signal_service_async.py @@ -0,0 +1,233 @@ +"""Tests for the async trade-lifecycle functions in +app/services/signal_service.py: `close_stale_trades` (the exit-rule engine — +time limit, stop loss, take profit, trailing stop) and `expire_old_signals`. + +Both functions open their own DB session via a module-level +`async_session_factory` rather than accepting one as a parameter, so these +tests monkeypatch that name to an in-memory SQLite session factory instead +of the real Postgres one. +""" +from __future__ import annotations + +import uuid +from datetime import datetime, timedelta, timezone +from decimal import Decimal + +import pytest +from sqlalchemy import select + +from app.models import Exchange, HypotheticalTrade, Signal, Symbol, User +from app.models.candle import Candle +from app.services import signal_service + + +def make_user(**prefs): + return User( + id=uuid.uuid4(), username=f"u-{uuid.uuid4().hex[:8]}", + email=f"{uuid.uuid4().hex[:8]}@example.com", password_hash="x", + role="trader", is_active=True, preferences=prefs, + ) + + +async def _seed_price(db, *, exchange_name: str, symbol: str, timeframe: str, price: Decimal): + exchange = Exchange(name=exchange_name, display_name=exchange_name.upper()) + db.add(exchange) + await db.flush() + sym = Symbol(symbol=symbol, base=symbol.split("/")[0], quote=symbol.split("/")[1], exchange_id=exchange.id) + db.add(sym) + await db.flush() + db.add(Candle( + symbol_id=sym.id, timeframe=timeframe, timestamp=datetime.now(timezone.utc), + open=price, high=price, low=price, close=price, volume=Decimal("1000"), + )) + await db.flush() + + +class TestCloseStaleTradesTimeLimit: + async def test_closes_trade_that_exceeds_max_hold_hours(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + user = make_user() + db.add(user) + await _seed_price(db, exchange_name="mexc", symbol="BTC/USDT", timeframe="1h", price=Decimal("100")) + trade = HypotheticalTrade( + user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=datetime.now(timezone.utc) - timedelta(hours=9), # exceeds default 8h + quantity=Decimal("1"), status="OPEN", + ) + db.add(trade) + await db.commit() + trade_id = trade.id + + await signal_service.close_stale_trades() + + async with session_factory() as db: + result = await db.execute(select(HypotheticalTrade).where(HypotheticalTrade.id == trade_id)) + saved = result.scalar_one() + assert saved.status == "CLOSED" + assert saved.exit_reason == "TIME_LIMIT" + + async def test_does_not_close_trade_within_hold_window_and_no_sl_tp_hit(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + user = make_user() + db.add(user) + # Small 2% move -- below default 5% SL and 10% TP + await _seed_price(db, exchange_name="mexc", symbol="BTC/USDT", timeframe="1h", price=Decimal("102")) + trade = HypotheticalTrade( + user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=datetime.now(timezone.utc) - timedelta(hours=1), + quantity=Decimal("1"), status="OPEN", + ) + db.add(trade) + await db.commit() + trade_id = trade.id + + await signal_service.close_stale_trades() + + async with session_factory() as db: + result = await db.execute(select(HypotheticalTrade).where(HypotheticalTrade.id == trade_id)) + saved = result.scalar_one() + assert saved.status == "OPEN" + + +class TestCloseStaleTradesStopLossTakeProfit: + async def test_closes_long_on_fixed_percent_stop_loss(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + user = make_user() + db.add(user) + # 6% drop -- exceeds default fixed stop_loss_pct=5% + await _seed_price(db, exchange_name="mexc", symbol="BTC/USDT", timeframe="1h", price=Decimal("94")) + trade = HypotheticalTrade( + user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=datetime.now(timezone.utc) - timedelta(hours=1), + quantity=Decimal("1"), status="OPEN", + ) + db.add(trade) + await db.commit() + trade_id = trade.id + + await signal_service.close_stale_trades() + + async with session_factory() as db: + result = await db.execute(select(HypotheticalTrade).where(HypotheticalTrade.id == trade_id)) + saved = result.scalar_one() + assert saved.status == "CLOSED" + assert saved.exit_reason == "STOP_LOSS" + assert saved.pnl == Decimal("-6") + + async def test_closes_long_on_fixed_percent_take_profit(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + user = make_user() + db.add(user) + # 12% rise -- exceeds default fixed take_profit_pct=10% + await _seed_price(db, exchange_name="mexc", symbol="BTC/USDT", timeframe="1h", price=Decimal("112")) + trade = HypotheticalTrade( + user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=datetime.now(timezone.utc) - timedelta(hours=1), + quantity=Decimal("1"), status="OPEN", + ) + db.add(trade) + await db.commit() + trade_id = trade.id + + await signal_service.close_stale_trades() + + async with session_factory() as db: + result = await db.execute(select(HypotheticalTrade).where(HypotheticalTrade.id == trade_id)) + saved = result.scalar_one() + assert saved.status == "CLOSED" + assert saved.exit_reason == "TARGET" + + +class TestCloseStaleTradesTrailingStop: + async def test_closes_when_price_drops_through_trailing_stop(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + user = make_user(auto_trade_trailing_stops={ + "BTC/USDT_mexc": { + "direction": "LONG", "best_price": 110.0, + "trailing_stop_price": 104.5, "trailing_pct": 5.0, + }, + }) + db.add(user) + # 4% move -- below the fixed 5% SL threshold, so SL doesn't + # preempt; but it's below the pre-set trailing_stop_price=104.5 + await _seed_price(db, exchange_name="mexc", symbol="BTC/USDT", timeframe="1h", price=Decimal("104")) + trade = HypotheticalTrade( + user_id=user.id, symbol="BTC/USDT", exchange="mexc", timeframe="1h", + direction="LONG", entry_price=Decimal("100"), + entry_time=datetime.now(timezone.utc) - timedelta(hours=1), + quantity=Decimal("1"), status="OPEN", + ) + db.add(trade) + await db.commit() + trade_id = trade.id + + await signal_service.close_stale_trades() + + async with session_factory() as db: + result = await db.execute(select(HypotheticalTrade).where(HypotheticalTrade.id == trade_id)) + saved = result.scalar_one() + assert saved.status == "CLOSED" + assert saved.exit_reason == "TRAILING_STOP" + + +class TestExpireOldSignals: + async def test_marks_old_active_signals_as_expired(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + old_signal = Signal( + symbol="BTC/USDT", exchange="mexc", timeframe="1h", + signal_type="BUY", strength="MODERATE", price=Decimal("100"), + timestamp=datetime.now(timezone.utc) - timedelta(days=10), + status="ACTIVE", + created_at=datetime.now(timezone.utc) - timedelta(days=10), + ) + recent_signal = Signal( + symbol="BTC/USDT", exchange="mexc", timeframe="1h", + signal_type="BUY", strength="MODERATE", price=Decimal("100"), + timestamp=datetime.now(timezone.utc) - timedelta(hours=1), + status="ACTIVE", + created_at=datetime.now(timezone.utc) - timedelta(hours=1), + ) + db.add_all([old_signal, recent_signal]) + await db.commit() + old_id, recent_id = old_signal.id, recent_signal.id + + expired_count = await signal_service.expire_old_signals(max_age_days=7) + assert expired_count == 1 + + async with session_factory() as db: + old_result = await db.execute(select(Signal).where(Signal.id == old_id)) + recent_result = await db.execute(select(Signal).where(Signal.id == recent_id)) + assert old_result.scalar_one().status == "EXPIRED" + assert recent_result.scalar_one().status == "ACTIVE" + + async def test_already_expired_signals_are_not_recounted(self, session_factory, monkeypatch): + monkeypatch.setattr(signal_service, "async_session_factory", session_factory) + + async with session_factory() as db: + db.add(Signal( + symbol="BTC/USDT", exchange="mexc", timeframe="1h", + signal_type="BUY", strength="MODERATE", price=Decimal("100"), + timestamp=datetime.now(timezone.utc) - timedelta(days=10), + status="EXPIRED", + created_at=datetime.now(timezone.utc) - timedelta(days=10), + )) + await db.commit() + + expired_count = await signal_service.expire_old_signals(max_age_days=7) + assert expired_count == 0