test: them 40 pytest cho indicator_service.py + phan async cua signal_service.py

- test_indicator_service.py (33 test): sma, ema, rsi, bollinger_bands,
  macd, atr, vwap, obv, obv_signal, mfi, detect_market_regime. Phat hien
  2 quirk nho (chua fix, can quyet dinh cua team):
    (q) rsi() tra ~98 thay vi 50 khi gia hoan toan di ngang (rs=50 sentinel
        van bi dua qua cong thuc RSI thay vi tra thang 50)
    (r) mfi() bi wraparound index o diem tinh dau tien cua chuoi (j-1=-1),
        tac dong thuc te gan bang 0 vi signal_service chi doc mfi_data[-1]
- test_signal_service_async.py (7 test): close_stale_trades (time limit,
  stop loss, take profit, trailing stop) + expire_old_signals. Cac ham
  nay tu mo session rieng qua async_session_factory (khong nhan db lam
  tham so) nen test monkeypatch bien module-level nay sang SQLite in-memory.
- conftest.py: them fixture session_factory (async_sessionmaker thay vi 1
  session) + _UTCDateTime TypeDecorator de SQLite giu duoc tzinfo UTC qua
  round-trip (SQLite khong ho tro luu tz-aware datetime nhu Postgres).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
Le
2026-07-03 22:30:31 +07:00
parent 9a0d2ab220
commit 06bc5ba29b
2 changed files with 510 additions and 0 deletions
+277
View File
@@ -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"
+233
View File
@@ -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