diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index bb75a53..d1e8987 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -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): diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 81e4e74..321c020 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -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) diff --git a/backend/app/api/signals.py b/backend/app/api/signals.py new file mode 100644 index 0000000..4edf6a9 --- /dev/null +++ b/backend/app/api/signals.py @@ -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) diff --git a/backend/app/application/services/signal_service.py b/backend/app/application/services/signal_service.py new file mode 100644 index 0000000..4c89219 --- /dev/null +++ b/backend/app/application/services/signal_service.py @@ -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) diff --git a/backend/app/domain/entities/signal.py b/backend/app/domain/entities/signal.py new file mode 100644 index 0000000..b74806f --- /dev/null +++ b/backend/app/domain/entities/signal.py @@ -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 diff --git a/backend/app/domain/repositories/signal.py b/backend/app/domain/repositories/signal.py new file mode 100644 index 0000000..f2680a9 --- /dev/null +++ b/backend/app/domain/repositories/signal.py @@ -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]: ... diff --git a/backend/app/infrastructure/persistence/migrations/versions/d8e0b2f3c4d5_signal_tables.py b/backend/app/infrastructure/persistence/migrations/versions/d8e0b2f3c4d5_signal_tables.py new file mode 100644 index 0000000..7db779b --- /dev/null +++ b/backend/app/infrastructure/persistence/migrations/versions/d8e0b2f3c4d5_signal_tables.py @@ -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") diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py index 5311611..e031465 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/__init__.py @@ -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, +) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/models/signal.py b/backend/app/infrastructure/persistence/sqlalchemy/models/signal.py new file mode 100644 index 0000000..ee441c8 --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/models/signal.py @@ -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) diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py new file mode 100644 index 0000000..2c31f6b --- /dev/null +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py @@ -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 diff --git a/backend/app/quant/signal.py b/backend/app/quant/signal.py new file mode 100644 index 0000000..02e08cb --- /dev/null +++ b/backend/app/quant/signal.py @@ -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")}, + ) diff --git a/backend/tests/test_signals.py b/backend/tests/test_signals.py new file mode 100644 index 0000000..fb1255b --- /dev/null +++ b/backend/tests/test_signals.py @@ -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 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