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 通过
This commit is contained in:
Simon
2026-09-09 00:35:37 +08:00
parent ef09d5b419
commit ba52edc2d6
12 changed files with 715 additions and 0 deletions
+18
View File
@@ -11,6 +11,7 @@ from fastapi import Depends
from sqlalchemy.orm import Session
from app.application.services.selection_service import SelectionService
from app.application.services.signal_service import SignalService
from app.domain.repositories.composite import CompositeRepository
from app.domain.repositories.factor import FactorRepository
from app.domain.repositories.jobs import ExperimentRepository, JobRepository
@@ -20,6 +21,7 @@ from app.domain.repositories.market import (
StockRepository,
)
from app.domain.repositories.selection import SelectionRepository
from app.domain.repositories.signal import SignalRepository
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
SqlAlchemyCompositeRepository,
)
@@ -34,6 +36,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
SqlAlchemySelectionRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.signal_impl import (
SqlAlchemySignalRepository,
)
from app.infrastructure.persistence.sqlalchemy.session import get_session
from app.quant.engine import LocalEngine, QuantEngine
from app.quant.service import ResearchService
@@ -65,6 +70,13 @@ def _service_factory(
return ResearchService(stock_repo, daily_repo, engine)
def _signal_service_factory(
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
) -> SignalService:
return SignalService(stock_repo, daily_repo)
def _selection_service_factory(
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
@@ -86,6 +98,10 @@ def _composite_repo_factory(session: DbSession) -> CompositeRepository:
return SqlAlchemyCompositeRepository(session)
def _signal_repo_factory(session: DbSession) -> SignalRepository:
return SqlAlchemySignalRepository(session)
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
EngineDep = Annotated[QuantEngine, Depends(_engine_factory)]
@@ -94,6 +110,8 @@ SelectionServiceDep = Annotated[SelectionService, Depends(_selection_service_fac
SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)]
FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_repo_factory)]
SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
def _job_repo_factory(session: DbSession):
+2
View File
@@ -17,6 +17,7 @@ from app.api import (
jobs,
research,
selections,
signals,
stocks,
)
@@ -27,6 +28,7 @@ api_router.include_router(factors.router)
api_router.include_router(composites.router)
api_router.include_router(research.router)
api_router.include_router(selections.router)
api_router.include_router(signals.router)
api_router.include_router(jobs.router)
api_router.include_router(experiments.router)
api_router.include_router(agent.router)
+65
View File
@@ -0,0 +1,65 @@
"""交易信号 API(M8.1):提交/查询信号(落库可复现)。
POST /api/signals body: {query: SelectionQuery, rules?: SignalRules}
GET /api/signals/{id} 读回某次信号
GET /api/signals 历史信号元数据(可过滤 as_of)
"""
from __future__ import annotations
from datetime import date
from typing import Annotated
from fastapi import APIRouter, HTTPException, Query
from pydantic import BaseModel
from app.api.deps import DbSession, SignalRepoDep, SignalServiceDep
from app.application.services.job_executor import new_id
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalMeta, SignalResult, SignalRules
router = APIRouter(prefix="/signals", tags=["signals"])
_AsOfQuery = Annotated[date | None, Query(description="按信号时点过滤")]
_LimitQuery = Annotated[int, Query(ge=1, le=200)]
class SignalRequest(BaseModel):
query: SelectionQuery
rules: SignalRules = SignalRules()
class SignalRun(BaseModel):
signal_id: str
result: SignalResult
@router.post("", response_model=SignalRun, summary="生成一次交易信号(同步)并落库")
def run_signal(
req: SignalRequest,
service: SignalServiceDep,
signal_repo: SignalRepoDep,
session: DbSession,
) -> SignalRun:
result = service.signal(req.query, req.rules)
signal_id = new_id("SIG")
signal_repo.save(signal_id, result)
session.commit()
return SignalRun(signal_id=signal_id, result=result)
@router.get("/{signal_id}", response_model=SignalResult, summary="读回一次信号结果")
def get_signal(signal_id: str, signal_repo: SignalRepoDep) -> SignalResult:
result = signal_repo.get(signal_id)
if result is None:
raise HTTPException(status_code=404, detail=f"信号记录 {signal_id} 不存在")
return result
@router.get("", response_model=list[SignalMeta], summary="历史信号元数据列表")
def list_signals(
signal_repo: SignalRepoDep,
as_of: _AsOfQuery = None,
limit: _LimitQuery = 20,
) -> list[SignalMeta]:
return signal_repo.list_recent(as_of=as_of, limit=limit)
@@ -0,0 +1,39 @@
"""信号用例(M8.1):基于选股评分排序 + 技术条件生成交易信号。
信号与回测买入逻辑同源(同一评分引擎、同一口径),保证「为什么 BUY/SELL」可解释。
"""
from __future__ import annotations
from datetime import date, timedelta
import pandas as pd
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import SignalResult, SignalRules
from app.domain.repositories.market import DailyBarRepository, StockRepository
from app.quant.selection import factor_columns
from app.quant.service import filter_stocks, load_daily_df
from app.quant.signal import generate_signals
class SignalService:
def __init__(self, stock_repo: StockRepository, daily_repo: DailyBarRepository) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
def signal(self, query: SelectionQuery, rules: SignalRules) -> SignalResult:
as_of = query.as_of or date.today()
stocks = filter_stocks(self._stock_repo.list(), query.universe, as_of=as_of)
if not stocks:
return generate_signals(pd.DataFrame(), query, rules, as_of)
columns = sorted(factor_columns(query))
daily = load_daily_df(
self._daily_repo,
[s.symbol for s in stocks],
as_of - timedelta(days=query.warmup_days),
as_of,
columns,
adjust=query.price_adjustment,
)
return generate_signals(daily, query, rules, as_of)
+57
View File
@@ -0,0 +1,57 @@
"""交易信号领域实体(M8.1,v2 §15 Signal Engine)。
Signal 输入 = Selection 排序(score)+ 价格/技术条件 + 规则;输出事件可解释:
BUY / WATCH / SELL(破位警示),每条带 trigger_reason —— 回答
「某日为什么对该股票给 BUY/SELL」(v2 §8)。
"""
from __future__ import annotations
from datetime import date, datetime
from pydantic import BaseModel, Field
class SignalRules(BaseModel):
"""规则(结构化,MVP):买入区间 + 趋势/动量条件 + 卖出/警示区间。"""
buy_rank_threshold: int = Field(default=20, ge=1, le=500, description="rank<=此值进入买入候选")
buy_require_trend: bool = Field(default=True, description="买入需 close > MA(trend_ma)")
buy_require_momentum: bool = Field(default=False, description="买入需 close > 20 日前 close")
trend_ma: int = Field(default=60, ge=10, le=250)
sell_rank_threshold: int = Field(default=50, ge=1, le=1000, description="rank>此值或破位 → SELL 警示")
sell_on_trend_break: bool = Field(default=True, description="买入区间内 close < MA(trend_ma) → SELL")
max_output_rank: int = Field(default=80, ge=1, le=2000, description="仅输出排名前 N 的信号")
class SignalEvent(BaseModel):
symbol: str
signal_date: date
signal_type: str = Field(pattern="^(BUY|WATCH|SELL)$")
score: float | None = None
price: float | None = None
trigger_reason: list[str] = Field(default_factory=list)
class SignalStatistics(BaseModel):
universe_size: int = 0
buy: int = 0
watch: int = 0
sell: int = 0
class SignalResult(BaseModel):
as_of_date: date
rules: SignalRules
statistics: SignalStatistics
events: list[SignalEvent] = Field(default_factory=list)
config_snapshot: dict = Field(default_factory=dict)
class SignalMeta(BaseModel):
id: str
as_of: date
buy: int = 0
watch: int = 0
sell: int = 0
created_at: datetime | None = None
+19
View File
@@ -0,0 +1,19 @@
"""Signal Repository Protocol(M8.1 落库)。"""
from __future__ import annotations
from datetime import date
from typing import Protocol
from app.domain.entities.signal import SignalMeta, SignalResult
class SignalRepository(Protocol):
def save(self, signal_id: str, result: SignalResult) -> None:
"""snapshot 一行 + 事件逐行(同事务,调用方 commit)。"""
def get(self, signal_id: str) -> SignalResult | None: ...
def list_recent(
self, as_of: date | None = None, limit: int = 20
) -> list[SignalMeta]: ...
@@ -0,0 +1,62 @@
"""signal_snapshot / signal_event 表(M8.1 交易信号)
Revision ID: d8e0b2f3c4d5
Revises: c3e9a0d1f4b5
Create Date: 2026-09-09
"""
from __future__ import annotations
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
revision: str = "d8e0b2f3c4d5"
down_revision: str | None = "c3e9a0d1f4b5"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
op.create_table(
"signal_snapshot",
sa.Column("id", sa.String(length=32), nullable=False),
sa.Column("as_of", sa.Date(), nullable=False),
sa.Column("query_json", sa.Text(), nullable=False),
sa.Column("rules_json", sa.Text(), nullable=False),
sa.Column("statistics_json", sa.Text(), nullable=False),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_signal_snapshot_as_of", "signal_snapshot", ["as_of"])
op.create_index("ix_signal_snapshot_created_at", "signal_snapshot", ["created_at"])
op.create_table(
"signal_event",
sa.Column(
"id",
sa.BigInteger().with_variant(sa.Integer(), "sqlite"),
autoincrement=True,
nullable=False,
),
sa.Column("signal_id", sa.String(length=32), nullable=False),
sa.Column("symbol", sa.String(length=12), nullable=False),
sa.Column("signal_date", sa.Date(), nullable=False),
sa.Column("signal_type", sa.String(length=8), nullable=False),
sa.Column("score", sa.Numeric(precision=14, scale=6), nullable=True),
sa.Column("price", sa.Numeric(precision=14, scale=4), nullable=True),
sa.Column("reason_json", sa.Text(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_signal_event_signal_id", "signal_event", ["signal_id"])
def downgrade() -> None:
op.drop_index("ix_signal_event_signal_id", table_name="signal_event")
op.drop_table("signal_event")
op.drop_index("ix_signal_snapshot_created_at", table_name="signal_snapshot")
op.drop_index("ix_signal_snapshot_as_of", table_name="signal_snapshot")
op.drop_table("signal_snapshot")
@@ -26,3 +26,7 @@ from app.infrastructure.persistence.sqlalchemy.models.selection import ( # noqa
SelectionResultModel,
SelectionSnapshotModel,
)
from app.infrastructure.persistence.sqlalchemy.models.signal import ( # noqa: F401
SignalEventModel,
SignalSnapshotModel,
)
@@ -0,0 +1,41 @@
"""交易信号表(M8.1)。
signal_snapshot:一次信号运行的查询/规则/统计快照
signal_event:逐信号(signal_type + score/price + trigger_reason JSON)
"""
from __future__ import annotations
from datetime import date, datetime
from sqlalchemy import BigInteger, Date, DateTime, Integer, Numeric, String, Text
from sqlalchemy.orm import Mapped, mapped_column
from app.infrastructure.persistence.sqlalchemy.base import Base
PK_INT = BigInteger().with_variant(Integer, "sqlite")
class SignalSnapshotModel(Base):
__tablename__ = "signal_snapshot"
id: Mapped[str] = mapped_column(String(32), primary_key=True)
as_of: Mapped[date] = mapped_column(Date, index=True)
query_json: Mapped[str] = mapped_column(Text)
rules_json: Mapped[str] = mapped_column(Text)
statistics_json: Mapped[str] = mapped_column(Text)
created_at: Mapped[datetime] = mapped_column(DateTime, index=True)
class SignalEventModel(Base):
__tablename__ = "signal_event"
id: Mapped[int] = mapped_column(PK_INT, primary_key=True, autoincrement=True)
signal_id: Mapped[str] = mapped_column(String(32), index=True)
symbol: Mapped[str] = mapped_column(String(12))
signal_date: Mapped[date] = mapped_column(Date)
signal_type: Mapped[str] = mapped_column(String(8))
score: Mapped[float | None] = mapped_column(Numeric(14, 6), nullable=True)
price: Mapped[float | None] = mapped_column(Numeric(14, 4), nullable=True)
reason_json: Mapped[str | None] = mapped_column(Text, nullable=True)
created_at: Mapped[datetime] = mapped_column(DateTime)
@@ -0,0 +1,108 @@
"""Signal Repository 的 SQLAlchemy 实现(M8.1)。"""
from __future__ import annotations
import json
from datetime import date, datetime
from sqlalchemy import select
from sqlalchemy.orm import Session
from app.domain.entities.signal import (
SignalEvent,
SignalMeta,
SignalResult,
SignalRules,
SignalStatistics,
)
from app.infrastructure.persistence.sqlalchemy.models.signal import (
SignalEventModel,
SignalSnapshotModel,
)
class SqlAlchemySignalRepository:
def __init__(self, session: Session) -> None:
self._session = session
def save(self, signal_id: str, result: SignalResult) -> None:
self._session.add(
SignalSnapshotModel(
id=signal_id,
as_of=result.as_of_date,
query_json=json.dumps(
{
"as_of": result.as_of_date.isoformat(),
**{k: v for k, v in result.config_snapshot.items() if k != "as_of"},
},
ensure_ascii=False,
),
rules_json=json.dumps(result.rules.model_dump(mode="json"), ensure_ascii=False),
statistics_json=json.dumps(result.statistics.model_dump(mode="json")),
created_at=datetime.now(),
)
)
now = datetime.now()
for e in result.events:
self._session.add(
SignalEventModel(
signal_id=signal_id,
symbol=e.symbol,
signal_date=e.signal_date,
signal_type=e.signal_type,
score=e.score,
price=e.price,
reason_json=json.dumps(e.trigger_reason, ensure_ascii=False),
created_at=now,
)
)
self._session.flush()
def get(self, signal_id: str) -> SignalResult | None:
snap = self._session.get(SignalSnapshotModel, signal_id)
if snap is None:
return None
rows = self._session.scalars(
select(SignalEventModel)
.where(SignalEventModel.signal_id == signal_id)
.order_by(SignalEventModel.signal_type, SignalEventModel.symbol)
).all()
rules = SignalRules.model_validate_json(snap.rules_json)
stats = SignalStatistics.model_validate_json(snap.statistics_json)
return SignalResult(
as_of_date=snap.as_of,
rules=rules,
statistics=stats,
events=[
SignalEvent(
symbol=r.symbol,
signal_date=r.signal_date,
signal_type=r.signal_type,
score=float(r.score) if r.score is not None else None,
price=float(r.price) if r.price is not None else None,
trigger_reason=json.loads(r.reason_json or "[]"),
)
for r in rows
],
config_snapshot={"as_of": snap.as_of.isoformat()},
)
def list_recent(self, as_of: date | None = None, limit: int = 20) -> list[SignalMeta]:
stmt = select(SignalSnapshotModel).order_by(SignalSnapshotModel.created_at.desc())
if as_of is not None:
stmt = stmt.where(SignalSnapshotModel.as_of == as_of)
stmt = stmt.limit(limit)
metas: list[SignalMeta] = []
for snap in self._session.scalars(stmt).all():
stats = SignalStatistics.model_validate_json(snap.statistics_json)
metas.append(
SignalMeta(
id=snap.id,
as_of=snap.as_of,
buy=stats.buy,
watch=stats.watch,
sell=stats.sell,
created_at=snap.created_at,
)
)
return metas
+135
View File
@@ -0,0 +1,135 @@
"""Signal Engine(v2 §15)—— 纯 pandas 执行。
输入:行情长表(<=as_of)+ SelectionQuery(评分因子)+ SignalRules。
流程:复合分 → 全市场 rank → 按规则判定 BUY / WATCH / SELL,
事件带 trigger_reason 与 score/price(可解释)。回测的买入逻辑(TopK+可买过滤)
与这里的 BUY 建议同源于同一评分引擎(v2 §25 一致)。
"""
from __future__ import annotations
from datetime import date
import pandas as pd
from app.domain.entities.selection import SelectionQuery
from app.domain.entities.signal import (
SignalEvent,
SignalResult,
SignalRules,
SignalStatistics,
)
from app.quant.composite import build_score_panel
from app.quant.selection import resolve_observation_date
def generate_signals(
daily: pd.DataFrame,
query: SelectionQuery,
rules: SignalRules,
as_of: date | None,
) -> SignalResult:
if not query.factors:
raise ValueError("signal 需要评分因子(SelectionQuery.factors)")
obs = resolve_observation_date(daily, as_of)
resolved = (obs.date() if obs is not None else as_of) or date.today()
if obs is None or daily.empty:
return _empty(query, rules, resolved)
view = daily[pd.to_datetime(daily["trade_date"]) <= obs]
score = build_score_panel(view, query.factors).loc[obs].dropna().sort_values(ascending=False)
close = view.pivot(index="trade_date", columns="symbol", values="close").sort_index()
close.index = pd.to_datetime(close.index)
c_d = close.loc[obs]
ma_d = close.rolling(rules.trend_ma).mean().loc[obs]
mom20_d = (close / close.shift(20) - 1.0).loc[obs]
events: list[SignalEvent] = []
stats = SignalStatistics(universe_size=int(len(score)))
for rank, (sym, sc) in enumerate(score.items(), start=1):
if rank > rules.max_output_rank:
break
c = _num(c_d.get(sym))
ma = _num(ma_d.get(sym))
mom = _num(mom20_d.get(sym))
price = float(c) if c is not None else None
reason: list[str] = []
event_type = "WATCH"
trend_ok = c is not None and ma is not None and c > ma
momentum_ok = mom is not None and mom > 0
if rank <= rules.buy_rank_threshold and (
not rules.buy_require_trend or trend_ok
) and (not rules.buy_require_momentum or momentum_ok):
event_type = "BUY"
reason = [f"综合分排名第 {rank}(≤买入阈值 {rules.buy_rank_threshold})"]
if rules.buy_require_trend:
reason.append(f"close > MA{rules.trend_ma}(趋势向上)")
if rules.buy_require_momentum:
reason.append("close > 20 日前收盘(动量为正)")
elif rank <= rules.buy_rank_threshold:
event_type = "WATCH"
reason = [f"综合分排名第 {rank}(买入区间)"]
if rules.buy_require_trend and not trend_ok:
reason.append(f"但 close < MA{rules.trend_ma}(趋势未确认)")
elif rank <= rules.sell_rank_threshold:
# 观望带
if rules.sell_on_trend_break and c is not None and ma is not None and c < ma:
event_type = "SELL"
reason = [f"跌破 MA{rules.trend_ma}(持仓者应卖出/减仓),rank={rank}"]
else:
event_type = "WATCH"
reason = [f"rank={rank}(买入区间外、卖出区间内:观望)"]
else:
event_type = "SELL"
reason = [f"综合分排名第 {rank}(>卖出阈值 {rules.sell_rank_threshold},持仓者应卖出)"]
if c is not None and c < 0:
continue # 防御负价
events.append(
SignalEvent(
symbol=sym,
signal_date=resolved,
signal_type=event_type,
score=round(float(sc), 6),
price=round(price, 4) if price is not None else None,
trigger_reason=reason,
)
)
if event_type == "BUY":
stats.buy += 1
elif event_type == "SELL":
stats.sell += 1
else:
stats.watch += 1
return SignalResult(
as_of_date=resolved,
rules=rules,
statistics=stats,
events=events,
config_snapshot={
"as_of": resolved.isoformat(),
"factors": [f.model_dump() for f in query.factors],
"rules": rules.model_dump(mode="json"),
},
)
def _num(v) -> float | None:
if v is None:
return None
try:
f = float(v)
except (TypeError, ValueError):
return None
return None if f != f else f # NaN → None
def _empty(query: SelectionQuery, rules: SignalRules, resolved: date) -> SignalResult:
return SignalResult(
as_of_date=resolved,
rules=rules,
statistics=SignalStatistics(),
events=[],
config_snapshot={"as_of": resolved.isoformat(), "rules": rules.model_dump(mode="json")},
)