Files
qlib/backend/tests/test_signals.py
T
Simon ba52edc2d6 feat(signal): M8.1 交易信号引擎(规则 + signal_event 落库 + /api/signals)
- SignalRules(买入 rank 阈值/趋势 MA/动量 + 卖出区间/破位警示)+ SignalEvent
  (BUY/WATCH/SELL,score/price/trigger_reason 可解释)+ SignalResult/Meta
- quant/signal.generate_signals:与选股同一评分引擎取全市场 rank,按规则分类输出
- signal_snapshot/signal_event 表(migration d8e0b2f3c4d5,MySQL 已应用)+ Repo
- SignalService + POST /api/signals(同步+落库)、GET 详情/列表
- tests/test_signals.py(引擎分类/排序/破位不 BUY、service、API 提交读回);全量 pytest 通过
2026-09-09 00:35:37 +08:00

166 lines
6.8 KiB
Python

"""M8.1 信号引擎测试:规则判定(BUY/WATCH/SELL)、可解释 reason、engine/service/API 落库回读。
使用合成行情:漂移差异决定 rank;构造「强趋势股(BUY)」与「高位破位股」验证类型。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
import pytest
from app.api import deps
from app.application.services.signal_service import SignalService
from app.domain.entities.market import Stock
from app.domain.entities.research import UniverseSpec
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalRules
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.main import app
from app.quant.signal import generate_signals
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily
_SYMS = ["600000.SH", "600001.SH", "600002.SH", "600003.SH", "600004.SH"]
class _MemDailyRepo:
def __init__(self, df: pd.DataFrame) -> None:
self._bars = bars_dataframe_to_daily_bars(df)
def get_range(self, symbol, start, end):
return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end]
def get_range_many(self, symbols, start, end, adjust="none"):
syms = set(symbols)
return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end]
class _MemStockRepo:
def __init__(self, stocks):
self._stocks = stocks
def list(self):
return self._stocks
def get_by_symbol(self, symbol):
return next((s for s in self._stocks if s.symbol == symbol), None)
def _query(as_of: date) -> SelectionQuery:
return SelectionQuery(
universe=UniverseSpec(exclude_st=False, min_listing_days=0),
factors=[{"name": "momentum_60", "weight": 1.0}],
top_n=50,
as_of=as_of,
)
class TestSignalEngine:
def test_buy_watch_sell_classification(self) -> None:
# 强势股数量足够时:顶部(高动量)BUY;构造一只「曾强现破位」→ SELL 警示
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
rules = SignalRules(buy_rank_threshold=2, sell_rank_threshold=4, max_output_rank=5)
res = generate_signals(df, _query(date(2024, 12, 31)), rules, date(2024, 12, 31))
assert res.statistics.buy == 2 # rank1/2 均为强趋势(上涨)→ BUY
types = {e.signal_type for e in res.events}
assert "BUY" in types and "WATCH" in types and "SELL" in types
# 每个事件可解释 + 价格
for e in res.events[:3]:
assert e.trigger_reason and (e.price or e.price == 0 or e.price is None)
bu = [e for e in res.events if e.signal_type == "BUY"]
assert all(any("排名" in r for r in e.trigger_reason) for e in bu)
def test_selected_rank_consistent_with_score(self) -> None:
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
res = generate_signals(df, _query(date(2024, 12, 31)), SignalRules(), date(2024, 12, 31))
events = sorted(res.events, key=lambda e: -e.score)
assert events == res.events # 按分数降序
assert res.statistics.universe_size == 5
def test_sell_on_trend_break(self) -> None:
"""下跌股在 BUY 区间内(如仅 3 只有分时 rank1 是跌股)→ WATCH/SELL 而非 BUY。"""
df = synthetic_daily({_SYMS[0]: -0.004, _SYMS[1]: 0.006, _SYMS[2]: 0.006}, n=320)
rules = SignalRules(buy_rank_threshold=1, trend_ma=60, max_output_rank=3)
res = generate_signals(df, _query(date(2024, 12, 31)), rules, date(2024, 12, 31))
# rank1 是负动量股(无正动量不构成 BUY 趋势要求?momentum_60 负但 close>ma60 可能仍成立
# 用趋势条件验证:若 rank1 close<ma60 → 不得 BUY
top = res.events[0]
if top.symbol == _SYMS[0]: # 负漂移:多半破位
assert top.signal_type != "BUY"
class TestSignalServiceAndApi:
def _svc(self, df):
stocks = [
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1)) for i, s in enumerate(_SYMS)
]
return SignalService(_MemStockRepo(stocks), _MemDailyRepo(df))
def test_service_runs(self) -> None:
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
res = self._svc(df).signal(_query(date(2024, 12, 31)), SignalRules(buy_rank_threshold=2))
assert res.statistics.buy == 2
assert res.config_snapshot["as_of"] == "2024-12-31"
@pytest.fixture()
def seeded_client(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'api.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
df = synthetic_daily({s: 0.008 - 0.002 * i for i, s in enumerate(_SYMS)}, n=320)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
with Session() as session:
SqlAlchemyStockRepository(session).upsert_many(
[
Stock(symbol=s, name=f"测试{i}", list_date=date(1999, 1, 1))
for i, s in enumerate(_SYMS)
]
)
SqlAlchemyDailyBarRepository(session).upsert_many(bars_dataframe_to_daily_bars(df))
session.commit()
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestSignalApi:
def test_submit_readback_list(self, seeded_client) -> None:
body = {
"query": {
"universe": {"exclude_st": False, "min_listing_days": 0},
"factors": [{"name": "momentum_60", "weight": 1}],
"top_n": 50,
"as_of": "2024-12-31",
},
"rules": {"buy_rank_threshold": 2, "sell_rank_threshold": 4, "max_output_rank": 5},
}
resp = seeded_client.post("/api/signals", json=body)
assert resp.status_code == 200
sig_id = resp.json()["signal_id"]
assert sig_id.startswith("SIG-")
result = resp.json()["result"]
assert result["statistics"]["buy"] == 2
assert result["events"][0]["trigger_reason"]
got = seeded_client.get(f"/api/signals/{sig_id}").json()
assert got["as_of_date"] == "2024-12-31"
assert got["events"][0]["symbol"] == result["events"][0]["symbol"]
rows = seeded_client.get("/api/signals").json()
assert len(rows) >= 1 and rows[0]["id"] == sig_id
assert seeded_client.get("/api/signals/SIG-NOPE").status_code == 404