feat(chart): M9-1 Chart DTO + Chart Service + Chart API(v3 §20)

- domain/entities/chart.py:ChartResult/OHLC/Volume/Series/EventMarker/ChartMetadata
  (adjust_mode + execution_price_basis 口径元数据)+ SelectionHit
- application/services/chart_service.py:个股 K线/量/MA 指标;显示层 qfq/hfq 折算
  (基于主口径 none 行情 × adjust_factor,绝回写研究数据);回测个股视图把实际成交
  转 fills 标记并在显示口径不同时做坐标换算(v3 §20.3/§20.5)
- by-symbol 历史查询:SignalRepository/SelectionRepository.list_by_symbol(含溯源 id)
- api/charts.py:/stocks/{symbol}/chart|signals|selections、/backtests/{id}/stocks/{symbol}/chart
  |trades|positions
- tests/test_charts.py(指标/qfq-hfq 折算断言/回测 fills/API 集成+404);全量 pytest 通过
This commit is contained in:
Simon
2026-09-09 07:09:52 +08:00
parent 8abfd6538c
commit 995ed08548
11 changed files with 736 additions and 1 deletions
+130
View File
@@ -0,0 +1,130 @@
"""Chart API(v3 §20.2):统一可视化数据只读接口。
- GET /api/stocks/{symbol}/chart?start&end&adjust K 线 + 量 + 指标(显示层折算)
- GET /api/stocks/{symbol}/signals 该股历史信号(markers)
- GET /api/stocks/{symbol}/selections 该股历史选股命中(markers)
- GET /api/backtests/{experiment_id}/stocks/{symbol}/chart 回测个股:K 线 + 实际成交 fills
- GET /api/backtests/{experiment_id}/trades|positions 回测成交/持仓展开
"""
from __future__ import annotations
from datetime import date
from typing import Annotated
from fastapi import APIRouter, HTTPException, Query
from app.api.deps import (
ChartServiceDep,
ExperimentRepoDep,
SelectionRepoDep,
SignalRepoDep,
)
from app.domain.entities.chart import ChartResult, EventMarker, SelectionHit
from app.domain.entities.research import BacktestResult
from app.domain.entities.signal import SignalHit
router = APIRouter(tags=["charts"])
_AdjustQuery = Annotated[str, Query(pattern="^(none|qfq|hfq)$")]
_StartQuery = Annotated[date | None, Query(description="开始日期(默认 2024-01-01)")]
_EndQuery = Annotated[date | None, Query(description="结束日期(默认今天)")]
def _signal_markers(hits: list[SignalHit]) -> list[EventMarker]:
kind_map = {"BUY": "signal_buy", "SELL": "signal_sell", "WATCH": "signal_watch"}
return [
EventMarker(
time=h.signal_date,
kind=kind_map.get(h.signal_type, "signal_watch"),
symbol="",
price=h.price,
score=h.score,
text=h.trigger_reason,
ref_id=h.signal_id,
)
for h in hits
]
def _selection_markers(hits: list[SelectionHit]) -> list[EventMarker]:
return [
EventMarker(
time=h.as_of,
kind="selection",
symbol=h.symbol,
score=h.score,
text=[f"rank #{h.rank}"] + h.selection_reason,
ref_id=h.selection_id,
)
for h in hits
]
@router.get("/stocks/{symbol}/chart", response_model=ChartResult, summary="个股 K 线图数据")
def stock_chart(
symbol: str,
service: ChartServiceDep,
start: _StartQuery = None,
end: _EndQuery = None,
adjust: _AdjustQuery = "none",
) -> ChartResult:
start = start or date(2024, 1, 1)
end = end or date.today()
if service.stock(symbol) is None:
raise HTTPException(status_code=404, detail=f"未找到股票 {symbol}")
return service.stock_chart(symbol, start, end, adjust)
@router.get("/stocks/{symbol}/signals", response_model=list[EventMarker], summary="该股历史信号")
def symbol_signals(symbol: str, signal_repo: SignalRepoDep) -> list[EventMarker]:
return _signal_markers(signal_repo.list_by_symbol(symbol))
@router.get("/stocks/{symbol}/selections", response_model=list[EventMarker], summary="该股历史选股命中")
def symbol_selections(symbol: str, selection_repo: SelectionRepoDep) -> list[EventMarker]:
return _selection_markers(selection_repo.list_by_symbol(symbol))
@router.get(
"/backtests/{experiment_id}/stocks/{symbol}/chart",
response_model=ChartResult,
summary="回测个股图(K 线 + 实际成交 fills)",
)
def backtest_stock_chart(
experiment_id: str,
symbol: str,
service: ChartServiceDep,
experiment_repo: ExperimentRepoDep,
start: _StartQuery = None,
end: _EndQuery = None,
adjust: _AdjustQuery = "none",
) -> ChartResult:
start = start or date(2024, 1, 1)
end = end or date.today()
exp = experiment_repo.get(experiment_id)
if exp is None:
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
try:
result = BacktestResult.model_validate_json(exp.result_json)
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=400, detail=f"{experiment_id} 不是 backtest 结果") from exc
return service.backtest_stock_chart(result, symbol, start, end, adjust)
@router.get("/backtests/{experiment_id}/trades", summary="回测成交明细")
def backtest_trades(experiment_id: str, experiment_repo: ExperimentRepoDep):
exp = experiment_repo.get(experiment_id)
if exp is None:
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
result = BacktestResult.model_validate_json(exp.result_json)
return result.trades
@router.get("/backtests/{experiment_id}/positions", summary="回测持仓明细")
def backtest_positions(experiment_id: str, experiment_repo: ExperimentRepoDep):
exp = experiment_repo.get(experiment_id)
if exp is None:
raise HTTPException(status_code=404, detail=f"Experiment {experiment_id} 不存在")
result = BacktestResult.model_validate_json(exp.result_json)
return result.positions
+16
View File
@@ -10,12 +10,14 @@ from typing import Annotated
from fastapi import Depends from fastapi import Depends
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.application.services.chart_service import ChartService
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.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
from app.domain.repositories.market import ( from app.domain.repositories.market import (
AdjustFactorRepository,
DailyBarRepository, DailyBarRepository,
FinancialRepository, FinancialRepository,
StockRepository, StockRepository,
@@ -30,6 +32,7 @@ from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import (
SqlAlchemyFactorRepository, SqlAlchemyFactorRepository,
) )
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyAdjustFactorRepository,
SqlAlchemyDailyBarRepository, SqlAlchemyDailyBarRepository,
SqlAlchemyFinancialRepository, SqlAlchemyFinancialRepository,
SqlAlchemyStockRepository, SqlAlchemyStockRepository,
@@ -62,6 +65,18 @@ def _financial_repo_factory(session: DbSession) -> FinancialRepository:
return SqlAlchemyFinancialRepository(session) return SqlAlchemyFinancialRepository(session)
def _adjust_repo_factory(session: DbSession) -> AdjustFactorRepository:
return SqlAlchemyAdjustFactorRepository(session)
def _chart_service_factory(
stock_repo: Annotated[StockRepository, Depends(_stock_repo_factory)],
daily_repo: Annotated[DailyBarRepository, Depends(_daily_repo_factory)],
adj_repo: Annotated[AdjustFactorRepository, Depends(_adjust_repo_factory)],
) -> ChartService:
return ChartService(stock_repo, daily_repo, adj_repo)
def _engine_factory() -> QuantEngine: def _engine_factory() -> QuantEngine:
return LocalEngine() return LocalEngine()
@@ -120,6 +135,7 @@ 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)] SignalRepoDep = Annotated[SignalRepository, Depends(_signal_repo_factory)]
SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)] SignalServiceDep = Annotated[SignalService, Depends(_signal_service_factory)]
ChartServiceDep = Annotated[ChartService, Depends(_chart_service_factory)]
StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)] StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)]
+2
View File
@@ -10,6 +10,7 @@ from fastapi import APIRouter
from app.api import ( from app.api import (
agent, agent,
charts,
composites, composites,
experiments, experiments,
factors, factors,
@@ -29,6 +30,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(charts.router)
api_router.include_router(signals.router) api_router.include_router(signals.router)
api_router.include_router(strategies.router) api_router.include_router(strategies.router)
api_router.include_router(jobs.router) api_router.include_router(jobs.router)
@@ -0,0 +1,206 @@
"""Chart Service(v3 §20.1)—— 只聚合与坐标整理,不重算研究结果。
- K 线/量/指标:基于主口径(adjust=none)行情,按请求 adjust 在**显示层**折算 qfq/hfq
(绝不回写研究数据;研究执行仍用 price_adjustment 指定口径)
- 标记:selections/signals 来自各自落库历史(by-symbol);backtest fills 来自 Experiment
内 BacktestResult 的 trades/positions
- 口径纪律(v3 §20.5):显示价 basis 与回测执行价 basis 分别记录;不一致时对 marker 做
与 K 线相同的坐标换算,保证成交点贴图
"""
from __future__ import annotations
from datetime import date
from app.domain.entities.chart import (
OHLC,
ChartMetadata,
ChartResult,
EventMarker,
SeriesPoint,
VolumePoint,
)
from app.domain.entities.market import AdjustFactor, DailyBar, Stock
from app.domain.entities.research import BacktestResult, Trade
from app.domain.repositories.market import (
AdjustFactorRepository,
DailyBarRepository,
StockRepository,
)
_MA_WINDOWS = (20, 60)
def _main_bars(bars: list[DailyBar]) -> list[DailyBar]:
"""仅取主口径行(adjust=none),丢弃新浪 qfq 兜底行 —— 显示层一律自 none 折算。"""
return [b for b in bars if b.adjust == "none"]
def _factor_multipliers(
adj_repo: AdjustFactorRepository,
symbol: str,
start: date,
end: date,
mode: str,
) -> dict[date, float]:
"""返回 {trade_date: 显示折算系数};mode=none → 空。qfq: f/f_latest;hfq: f。"""
if mode == "none":
return {}
factors: list[AdjustFactor] = adj_repo.get_range(symbol, date(1990, 1, 1), end)
if not factors:
return {}
by_day = {f.trade_date: float(f.factor) for f in factors}
latest = max(by_day.values())
out: dict[date, float] = {}
for day, f in by_day.items():
out[day] = f / latest if mode == "qfq" else f
return out
class ChartService:
def __init__(
self,
stock_repo: StockRepository,
daily_repo: DailyBarRepository,
adj_repo: AdjustFactorRepository,
) -> None:
self._stock_repo = stock_repo
self._daily_repo = daily_repo
self._adj_repo = adj_repo
def stock(self, symbol: str) -> Stock | None:
return self._stock_repo.get_by_symbol(symbol)
def stock_chart(
self,
symbol: str,
start: date,
end: date,
adjust: str = "none",
execution_price_basis: str | None = None,
extra_markers: list[EventMarker] | None = None,
) -> ChartResult:
"""基础个股 K 线图(可叠加 fills 等外部标记)。"""
stock = self.stock(symbol)
name = stock.name if stock else ""
raw = _main_bars(self._daily_repo.get_range(symbol, start, end))
mult = _factor_multipliers(self._adj_repo, symbol, start, end, adjust)
bars: list[OHLC] = []
volume: list[VolumePoint] = []
for b in raw:
m = mult.get(b.trade_date, 1.0)
bars.append(
OHLC(
time=b.trade_date,
open=_v(b.open, m),
high=_v(b.high, m),
low=_v(b.low, m),
close=_v(b.close, m),
)
)
volume.append(VolumePoint(time=b.trade_date, value=_v(b.volume, 1.0)))
markers = _convert_markers(extra_markers or [], mult)
indicators = _ma_indicators(bars)
return ChartResult(
metadata=ChartMetadata(
symbol=symbol,
name=name,
adjust_mode=adjust,
execution_price_basis=execution_price_basis,
start=start,
end=end,
bar_count=len(bars),
indicator_windows=list(_MA_WINDOWS),
),
bars=bars,
volume=volume,
indicators=indicators,
fills=[m for m in markers if m.kind.startswith("fill_")],
signals=[m for m in markers if m.kind.startswith("signal_")],
selections=[m for m in markers if m.kind == "selection"],
holding_periods=_holding_periods(bars),
)
# ---- 由已存历史构造标记(不重算) ----
def backtest_stock_chart(
self,
result: BacktestResult,
symbol: str,
start: date,
end: date,
adjust: str = "none",
) -> ChartResult:
"""回测个股视图:K 线 + 该股实际成交 fills(v3 §20.3 Signal↔Fill 展示)。"""
basis = (result.config_snapshot or {}).get("price_adjustment", "none")
markers = _trades_to_markers(result.trades, symbol, basis)
return self.stock_chart(symbol, start, end, adjust, execution_price_basis=basis,
extra_markers=markers)
def _trades_to_markers(trades: list[Trade], symbol: str, basis: str) -> list[EventMarker]:
markers: list[EventMarker] = []
for t in trades:
if t.symbol != symbol:
continue
markers.append(
EventMarker(
time=t.entry_date,
kind="fill_buy",
symbol=symbol,
price=t.entry_price,
text=[f"买入 @ {t.entry_price:.2f}(basis={basis})"],
)
)
markers.append(
EventMarker(
time=t.exit_date,
kind="fill_sell",
symbol=symbol,
price=t.exit_price,
text=[f"卖出 @ {t.exit_price:.2f},收益 {t.return_pct:.2f}%(basis={basis})"],
)
)
return markers
def _convert_markers(markers: list[EventMarker], mult: dict[date, float]) -> list[EventMarker]:
"""显示口径与执行价 basis 不一致时,把 marker 价格折算到 K 线坐标系(v3 §20.5)。"""
out: list[EventMarker] = []
for m in markers:
if m.price is not None and mult:
k = mult.get(m.time)
if k is not None:
m = m.model_copy(update={"price": round(m.price * k, 4)})
out.append(m)
return out
def _holding_periods(bars: list[OHLC]) -> list[dict]:
"""v1 空实现占位(持仓区间渲染 v3 §20.4 后续细化)。"""
return []
def _ma_indicators(bars: list[OHLC]) -> dict[str, list[SeriesPoint]]:
import statistics
closes = [b.close for b in bars]
out: dict[str, list[SeriesPoint]] = {}
for w in _MA_WINDOWS:
series: list[SeriesPoint] = []
for i, b in enumerate(bars):
if i + 1 < w:
continue
window = closes[i + 1 - w : i + 1]
if all(v is not None for v in window):
series.append(SeriesPoint(time=b.time, value=round(statistics.fmean(window), 4)))
out[f"ma{w}"] = series
return out
def _v(v, m: float) -> float | None:
if v is None:
return None
f = float(v)
return round(f * m, 4)
+88
View File
@@ -0,0 +1,88 @@
"""统一量化可视化 DTO(v3 §20 Chart Service)。
原则:前端只展示本结构,不得自行重算选股/信号/成交(v3 §20.1)。
价格口径:bars 已按请求 adjust 折算(显示层);成交/信号标记与 K 线同坐标系;
metadata 记录 adjust_mode 与回测执行价 basis,杜绝图表与回测口径混用(v3 §20.5)。
"""
from __future__ import annotations
from datetime import date
from pydantic import BaseModel, Field
class OHLC(BaseModel):
time: date
open: float | None = None
high: float | None = None
low: float | None = None
close: float | None = None
class VolumePoint(BaseModel):
time: date
value: float | None = None
class SeriesPoint(BaseModel):
"""指标/分数等 (时间, 值) 序列点。"""
time: date
value: float | None = None
class EventMarker(BaseModel):
"""K 线上可点击的事件标记(选股/信号/实际成交)。"""
time: date
kind: str = Field(
description="selection | signal_buy | signal_sell | signal_watch | fill_buy | fill_sell"
)
symbol: str = ""
price: float | None = None
score: float | None = None
text: list[str] = Field(default_factory=list, description="原因/说明(tooltip)")
ref_id: str | None = Field(default=None, description="关联 selection/signal 记录 id")
class ChartMetadata(BaseModel):
symbol: str
name: str = ""
adjust_mode: str = Field(default="none", description="显示口径:none | qfq | hfq")
price_basis: str = Field(default="chart_display", description="显示价基准(显示层折算)")
execution_price_basis: str | None = Field(
default=None, description="回测执行价口径(如 none),与显示口径不同时用于解释"
)
start: date | None = None
end: date | None = None
bar_count: int = 0
indicator_windows: list[int] = Field(default_factory=list)
class ChartResult(BaseModel):
"""单只股票 / 回测个股的统一图表数据(v3 §20.2)。"""
metadata: ChartMetadata
bars: list[OHLC] = Field(default_factory=list)
volume: list[VolumePoint] = Field(default_factory=list)
indicators: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
selections: list[EventMarker] = Field(default_factory=list)
signals: list[EventMarker] = Field(default_factory=list)
fills: list[EventMarker] = Field(default_factory=list)
holding_periods: list[dict] = Field(default_factory=list)
factor_values: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
strategy_scores: dict[str, list[SeriesPoint]] = Field(default_factory=dict)
unimplemented: list[str] = Field(default_factory=list)
class SelectionHit(BaseModel):
"""个股在选股历史中的命中(by-symbol 查询)。"""
selection_id: str
as_of: date
method: str
symbol: str
rank: int
score: float
selection_reason: list[str] = Field(default_factory=list)
+11
View File
@@ -55,3 +55,14 @@ class SignalMeta(BaseModel):
watch: int = 0 watch: int = 0
sell: int = 0 sell: int = 0
created_at: datetime | None = None created_at: datetime | None = None
class SignalHit(BaseModel):
"""个股在信号历史中的命中(by-symbol 查询,供 Chart 标记)。"""
signal_id: str
signal_date: date
signal_type: str
score: float | None = None
price: float | None = None
trigger_reason: list[str] = Field(default_factory=list)
@@ -10,6 +10,7 @@ from __future__ import annotations
from datetime import date from datetime import date
from typing import Protocol from typing import Protocol
from app.domain.entities.chart import SelectionHit
from app.domain.entities.selection import SelectionMeta, SelectionResult from app.domain.entities.selection import SelectionMeta, SelectionResult
@@ -27,3 +28,6 @@ class SelectionRepository(Protocol):
limit: int = 20, limit: int = 20,
) -> list[SelectionMeta]: ) -> list[SelectionMeta]:
"""历史选股元数据列表(按 created_at 倒序;可选 as_of/method 过滤)。""" """历史选股元数据列表(按 created_at 倒序;可选 as_of/method 过滤)。"""
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SelectionHit]:
"""该股在历史选股中的命中(Chart 标记用,含 as_of/rank/reason)。"""
+4 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
from datetime import date from datetime import date
from typing import Protocol from typing import Protocol
from app.domain.entities.signal import SignalMeta, SignalResult from app.domain.entities.signal import SignalHit, SignalMeta, SignalResult
class SignalRepository(Protocol): class SignalRepository(Protocol):
@@ -17,3 +17,6 @@ class SignalRepository(Protocol):
def list_recent( def list_recent(
self, as_of: date | None = None, limit: int = 20 self, as_of: date | None = None, limit: int = 20
) -> list[SignalMeta]: ... ) -> list[SignalMeta]: ...
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SignalHit]:
"""该股在历史信号中的命中(Chart 标记用,含 signal_id 溯源)。"""
@@ -12,6 +12,7 @@ from datetime import date, datetime
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from app.domain.entities.chart import SelectionHit
from app.domain.entities.selection import ( from app.domain.entities.selection import (
SelectionCandidate, SelectionCandidate,
SelectionMeta, SelectionMeta,
@@ -84,6 +85,27 @@ class SqlAlchemySelectionRepository:
config_snapshot=json.loads(snap.query_json), config_snapshot=json.loads(snap.query_json),
) )
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SelectionHit]:
rows = self._session.execute(
select(SelectionResultModel, SelectionSnapshotModel.as_of, SelectionSnapshotModel.method)
.join(SelectionSnapshotModel, SelectionSnapshotModel.id == SelectionResultModel.selection_id)
.where(SelectionResultModel.symbol == symbol)
.order_by(SelectionSnapshotModel.as_of.desc(), SelectionResultModel.rank)
.limit(limit)
).all()
return [
SelectionHit(
selection_id=r[0].selection_id,
as_of=r[1],
method=r[2],
symbol=r[0].symbol,
rank=r[0].rank,
score=float(r[0].score),
selection_reason=json.loads(r[0].reason_json or "[]"),
)
for r in rows
]
def list_recent( def list_recent(
self, self,
as_of: date | None = None, as_of: date | None = None,
@@ -10,6 +10,7 @@ from sqlalchemy.orm import Session
from app.domain.entities.signal import ( from app.domain.entities.signal import (
SignalEvent, SignalEvent,
SignalHit,
SignalMeta, SignalMeta,
SignalResult, SignalResult,
SignalRules, SignalRules,
@@ -87,6 +88,25 @@ class SqlAlchemySignalRepository:
config_snapshot={"as_of": snap.as_of.isoformat()}, config_snapshot={"as_of": snap.as_of.isoformat()},
) )
def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SignalHit]:
rows = self._session.scalars(
select(SignalEventModel)
.where(SignalEventModel.symbol == symbol)
.order_by(SignalEventModel.signal_date.desc(), SignalEventModel.id.desc())
.limit(limit)
).all()
return [
SignalHit(
signal_id=r.signal_id,
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
]
def list_recent(self, as_of: date | None = None, limit: int = 20) -> list[SignalMeta]: def list_recent(self, as_of: date | None = None, limit: int = 20) -> list[SignalMeta]:
stmt = select(SignalSnapshotModel).order_by(SignalSnapshotModel.created_at.desc()) stmt = select(SignalSnapshotModel).order_by(SignalSnapshotModel.created_at.desc())
if as_of is not None: if as_of is not None:
+233
View File
@@ -0,0 +1,233 @@
"""M9-1 Chart Service/API 测试:K线/量/指标、复权显示折算、按股历史信号与选股、
回测个股图(fills)、API 集成。使用 tmp SQLite + 真实 SQLAlchemy repo。
"""
from __future__ import annotations
from datetime import date, datetime
from decimal import Decimal
import pandas as pd
import pytest
from app.api import deps
from app.application.services.chart_service import ChartService
from app.domain.entities.market import AdjustFactor, Stock
from app.domain.entities.research import (
ExperimentRecord,
ResearchSpec,
)
from app.infrastructure.persistence.sqlalchemy.base import Base
from app.infrastructure.persistence.sqlalchemy.repositories.jobs_impl import (
SqlAlchemyExperimentRepository,
)
from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import (
SqlAlchemyAdjustFactorRepository,
SqlAlchemyDailyBarRepository,
SqlAlchemyStockRepository,
)
from app.main import app
from app.quant.engine import LocalEngine
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"]
def _daily_df(n=320) -> pd.DataFrame:
return synthetic_daily({s: 0.004 - 0.001 * i for i, s in enumerate(_SYMS)}, n=n)
@pytest.fixture()
def seeded(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'chart.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
df = _daily_df()
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()
return engine, Session, df
class TestChartServiceUnit:
def test_stock_chart_ohlc_and_indicators(self, seeded) -> None:
engine, Session, _df = seeded
with Session() as session:
svc = ChartService(
SqlAlchemyStockRepository(session),
SqlAlchemyDailyBarRepository(session),
SqlAlchemyAdjustFactorRepository(session),
)
res = svc.stock_chart(_SYMS[0], date(2024, 1, 1), date(2024, 12, 31), "none")
assert res.metadata.symbol == _SYMS[0]
assert res.metadata.bar_count > 200
assert res.bars[0].close and res.bars[-1].close
assert "ma20" in res.indicators and "ma60" in res.indicators
assert res.metadata.adjust_mode == "none"
def test_qfq_display_conversion(self, seeded, tmp_path) -> None:
"""前段因子 1.0 / 后段 2.0:qfq 显示把前段折半、hfq 前段不变后段翻倍。"""
engine, Session, df = _seeded_with_factors(tmp_path)
dates = sorted(df["trade_date"].unique())
split = dates[len(dates) // 2]
with Session() as session:
repo = SqlAlchemyAdjustFactorRepository(session)
rows = [
AdjustFactor(symbol=_SYMS[0], trade_date=d,
factor=Decimal("1.0") if d < split else Decimal("2.0"))
for d in dates
]
repo.upsert_many(rows)
session.commit()
svc = ChartService(
SqlAlchemyStockRepository(session),
SqlAlchemyDailyBarRepository(session),
SqlAlchemyAdjustFactorRepository(session),
)
none_chart = svc.stock_chart(_SYMS[0], dates[0], dates[-1], "none")
none_c = none_chart.bars[0].close
none_last = none_chart.bars[-1].close
qfq_c = svc.stock_chart(_SYMS[0], dates[0], dates[-1], "qfq").bars[0].close
hfq_c = svc.stock_chart(_SYMS[0], dates[0], dates[-1], "hfq").bars[0].close
last_qfq = svc.stock_chart(_SYMS[0], dates[0], dates[-1], "qfq").bars[-1].close
assert none_c is not None and qfq_c is not None and none_last is not None
assert abs(none_c - qfq_c * 2) < 1e-3 # qfq 前段折半
assert abs(none_c - hfq_c) < 1e-3 # hfq 前段不变
assert abs(none_last - last_qfq) < 1e-3 # 最新段 qfq 基准=原价
def test_backtest_chart_fills(self, seeded) -> None:
engine, Session, df = seeded
spec = ResearchSpec(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0,
"symbols": [_SYMS[0], _SYMS[1], _SYMS[2]]},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 12, 31)),
)
result = LocalEngine().run_backtest(df, spec)
with Session() as session:
svc = ChartService(
SqlAlchemyStockRepository(session),
SqlAlchemyDailyBarRepository(session),
SqlAlchemyAdjustFactorRepository(session),
)
chart = svc.backtest_stock_chart(result, _SYMS[0], date(2024, 1, 1), date(2024, 12, 31))
assert chart.metadata.execution_price_basis == "none"
assert chart.bars and chart.metadata.bar_count > 100
kinds = {f.kind for f in chart.fills}
assert kinds <= {"fill_buy", "fill_sell"}
# 有成交则有 fill 标记
if result.trades:
assert len(chart.fills) >= 2
def _seeded_with_factors(tmp_path):
engine = create_engine(f"sqlite:///{tmp_path / 'chart2.db'}", future=True)
Base.metadata.create_all(engine)
Session = sessionmaker(bind=engine, expire_on_commit=False)
df = _daily_df()
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()
return engine, Session, df
@pytest.fixture()
def client(tmp_path):
engine, Session, df = _seeded_with_factors(tmp_path)
def _session_override():
with Session() as s:
yield s
app.dependency_overrides[deps.get_session] = _session_override
# 预置一只实验(backtest),供 backtest chart API
spec = ResearchSpec(
type="backtest",
universe={"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS[:3]},
factors=[{"name": "momentum_60", "weight": 1.0}],
selection={"top_n": 2},
rebalance="monthly",
period=(date(2024, 5, 1), date(2024, 12, 31)),
)
result = LocalEngine().run_backtest(df, spec)
with Session() as session:
repo = SqlAlchemyExperimentRepository(session)
repo.save(
ExperimentRecord(
id="EXP-CHART-1",
kind="backtest",
spec_json=spec.model_dump_json(),
result_json=result.model_dump_json(),
summary_text="chart-test",
created_at=datetime.now(),
)
)
session.commit()
with TestClient(app) as c:
yield c
app.dependency_overrides.clear()
class TestChartApi:
def test_stock_chart_200(self, client) -> None:
resp = client.get(
f"/api/stocks/{_SYMS[0]}/chart?start=2024-01-01&end=2024-12-31&adjust=none"
)
assert resp.status_code == 200
body = resp.json()
assert body["metadata"]["adjust_mode"] == "none"
assert len(body["bars"]) > 200
assert "ma20" in body["indicators"]
def test_signals_selections_by_symbol(self, client) -> None:
# 生成一次信号 + 一次选股,再按 symbol 查询
body = {
"query": {
"universe": {"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
"factors": [{"name": "momentum_60", "weight": 1}],
"top_n": 10,
"as_of": "2024-12-31",
}
}
assert client.post("/api/signals", json={**body, "rules": {"buy_rank_threshold": 1, "max_output_rank": 4}}).status_code == 200
assert client.post("/api/selections", json={
"universe": {"exclude_st": False, "min_listing_days": 0, "symbols": _SYMS},
"method": "score", "factors": [{"name": "momentum_60", "weight": 1}],
"top_n": 2, "as_of": "2024-12-31",
}).status_code == 200
markers = client.get(f"/api/stocks/{_SYMS[0]}/signals").json()
assert isinstance(markers, list) and len(markers) >= 1
assert markers[0]["kind"].startswith("signal_")
assert markers[0]["text"]
hits = client.get(f"/api/stocks/{_SYMS[0]}/selections").json()
assert isinstance(hits, list)
assert all(m["kind"] == "selection" for m in hits)
def test_backtest_chart_and_trades(self, client) -> None:
chart = client.get(
f"/api/backtests/EXP-CHART-1/stocks/{_SYMS[0]}/chart?start=2024-01-01&end=2024-12-31"
)
assert chart.status_code == 200
c = chart.json()
assert c["metadata"]["execution_price_basis"] == "none"
assert c["bars"] and c["metadata"]["bar_count"] > 100
trades = client.get("/api/backtests/EXP-CHART-1/trades").json()
positions = client.get("/api/backtests/EXP-CHART-1/positions").json()
assert isinstance(trades, list) and isinstance(positions, list)
assert client.get("/api/backtests/EXP-NOPE/stocks/x/chart").status_code == 404