- 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 通过
109 lines
3.8 KiB
Python
109 lines
3.8 KiB
Python
"""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
|