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:
@@ -11,6 +11,7 @@ from fastapi import Depends
|
|||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from app.application.services.selection_service import SelectionService
|
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.composite import CompositeRepository
|
||||||
from app.domain.repositories.factor import FactorRepository
|
from app.domain.repositories.factor import FactorRepository
|
||||||
from app.domain.repositories.jobs import ExperimentRepository, JobRepository
|
from app.domain.repositories.jobs import ExperimentRepository, JobRepository
|
||||||
@@ -20,6 +21,7 @@ from app.domain.repositories.market import (
|
|||||||
StockRepository,
|
StockRepository,
|
||||||
)
|
)
|
||||||
from app.domain.repositories.selection import SelectionRepository
|
from app.domain.repositories.selection import SelectionRepository
|
||||||
|
from app.domain.repositories.signal import SignalRepository
|
||||||
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
from app.infrastructure.persistence.sqlalchemy.repositories.composite_impl import (
|
||||||
SqlAlchemyCompositeRepository,
|
SqlAlchemyCompositeRepository,
|
||||||
)
|
)
|
||||||
@@ -34,6 +36,9 @@ from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
|
|||||||
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
|
from app.infrastructure.persistence.sqlalchemy.repositories.selection_impl import (
|
||||||
SqlAlchemySelectionRepository,
|
SqlAlchemySelectionRepository,
|
||||||
)
|
)
|
||||||
|
from app.infrastructure.persistence.sqlalchemy.repositories.signal_impl import (
|
||||||
|
SqlAlchemySignalRepository,
|
||||||
|
)
|
||||||
from app.infrastructure.persistence.sqlalchemy.session import get_session
|
from app.infrastructure.persistence.sqlalchemy.session import get_session
|
||||||
from app.quant.engine import LocalEngine, QuantEngine
|
from app.quant.engine import LocalEngine, QuantEngine
|
||||||
from app.quant.service import ResearchService
|
from app.quant.service import ResearchService
|
||||||
@@ -65,6 +70,13 @@ def _service_factory(
|
|||||||
return ResearchService(stock_repo, daily_repo, engine)
|
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(
|
def _selection_service_factory(
|
||||||
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
|
||||||
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
|
||||||
@@ -86,6 +98,10 @@ def _composite_repo_factory(session: DbSession) -> CompositeRepository:
|
|||||||
return SqlAlchemyCompositeRepository(session)
|
return SqlAlchemyCompositeRepository(session)
|
||||||
|
|
||||||
|
|
||||||
|
def _signal_repo_factory(session: DbSession) -> SignalRepository:
|
||||||
|
return SqlAlchemySignalRepository(session)
|
||||||
|
|
||||||
|
|
||||||
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
StockRepoDep = Annotated[StockRepository, Depends(_stock_repo_factory)]
|
||||||
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
DailyRepoDep = Annotated[DailyBarRepository, Depends(_daily_repo_factory)]
|
||||||
EngineDep = Annotated[QuantEngine, Depends(_engine_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)]
|
SelectionRepoDep = Annotated[SelectionRepository, Depends(_selection_repo_factory)]
|
||||||
FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
|
FactorRepoDep = Annotated[FactorRepository, Depends(_factor_repo_factory)]
|
||||||
CompositeRepoDep = Annotated[CompositeRepository, Depends(_composite_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):
|
def _job_repo_factory(session: DbSession):
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from app.api import (
|
|||||||
jobs,
|
jobs,
|
||||||
research,
|
research,
|
||||||
selections,
|
selections,
|
||||||
|
signals,
|
||||||
stocks,
|
stocks,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -27,6 +28,7 @@ api_router.include_router(factors.router)
|
|||||||
api_router.include_router(composites.router)
|
api_router.include_router(composites.router)
|
||||||
api_router.include_router(research.router)
|
api_router.include_router(research.router)
|
||||||
api_router.include_router(selections.router)
|
api_router.include_router(selections.router)
|
||||||
|
api_router.include_router(signals.router)
|
||||||
api_router.include_router(jobs.router)
|
api_router.include_router(jobs.router)
|
||||||
api_router.include_router(experiments.router)
|
api_router.include_router(experiments.router)
|
||||||
api_router.include_router(agent.router)
|
api_router.include_router(agent.router)
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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]: ...
|
||||||
+62
@@ -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,
|
SelectionResultModel,
|
||||||
SelectionSnapshotModel,
|
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
|
||||||
@@ -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")},
|
||||||
|
)
|
||||||
@@ -0,0 +1,165 @@
|
|||||||
|
"""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
|
||||||
Reference in New Issue
Block a user