diff --git a/backend/app/api/charts.py b/backend/app/api/charts.py new file mode 100644 index 0000000..85138f0 --- /dev/null +++ b/backend/app/api/charts.py @@ -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 diff --git a/backend/app/api/deps.py b/backend/app/api/deps.py index df6f264..a61402f 100644 --- a/backend/app/api/deps.py +++ b/backend/app/api/deps.py @@ -10,12 +10,14 @@ from typing import Annotated from fastapi import Depends 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.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 from app.domain.repositories.market import ( + AdjustFactorRepository, DailyBarRepository, FinancialRepository, StockRepository, @@ -30,6 +32,7 @@ from app.infrastructure.persistence.sqlalchemy.repositories.factor_impl import ( SqlAlchemyFactorRepository, ) from app.infrastructure.persistence.sqlalchemy.repositories.market_impl import ( + SqlAlchemyAdjustFactorRepository, SqlAlchemyDailyBarRepository, SqlAlchemyFinancialRepository, SqlAlchemyStockRepository, @@ -62,6 +65,18 @@ def _financial_repo_factory(session: DbSession) -> FinancialRepository: 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: return LocalEngine() @@ -120,6 +135,7 @@ 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)] +ChartServiceDep = Annotated[ChartService, Depends(_chart_service_factory)] StrategyRepoDep = Annotated[StrategyRepository, Depends(_strategy_repo_factory)] diff --git a/backend/app/api/router.py b/backend/app/api/router.py index 23469eb..32d3f97 100644 --- a/backend/app/api/router.py +++ b/backend/app/api/router.py @@ -10,6 +10,7 @@ from fastapi import APIRouter from app.api import ( agent, + charts, composites, experiments, factors, @@ -29,6 +30,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(charts.router) api_router.include_router(signals.router) api_router.include_router(strategies.router) api_router.include_router(jobs.router) diff --git a/backend/app/application/services/chart_service.py b/backend/app/application/services/chart_service.py new file mode 100644 index 0000000..0988759 --- /dev/null +++ b/backend/app/application/services/chart_service.py @@ -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) diff --git a/backend/app/domain/entities/chart.py b/backend/app/domain/entities/chart.py new file mode 100644 index 0000000..46e8a73 --- /dev/null +++ b/backend/app/domain/entities/chart.py @@ -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) diff --git a/backend/app/domain/entities/signal.py b/backend/app/domain/entities/signal.py index b74806f..47fe47f 100644 --- a/backend/app/domain/entities/signal.py +++ b/backend/app/domain/entities/signal.py @@ -55,3 +55,14 @@ class SignalMeta(BaseModel): watch: int = 0 sell: int = 0 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) diff --git a/backend/app/domain/repositories/selection.py b/backend/app/domain/repositories/selection.py index 6d3a201..61600e8 100644 --- a/backend/app/domain/repositories/selection.py +++ b/backend/app/domain/repositories/selection.py @@ -10,6 +10,7 @@ from __future__ import annotations from datetime import date from typing import Protocol +from app.domain.entities.chart import SelectionHit from app.domain.entities.selection import SelectionMeta, SelectionResult @@ -27,3 +28,6 @@ class SelectionRepository(Protocol): limit: int = 20, ) -> list[SelectionMeta]: """历史选股元数据列表(按 created_at 倒序;可选 as_of/method 过滤)。""" + + def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SelectionHit]: + """该股在历史选股中的命中(Chart 标记用,含 as_of/rank/reason)。""" diff --git a/backend/app/domain/repositories/signal.py b/backend/app/domain/repositories/signal.py index f2680a9..c43d247 100644 --- a/backend/app/domain/repositories/signal.py +++ b/backend/app/domain/repositories/signal.py @@ -5,7 +5,7 @@ from __future__ import annotations from datetime import date from typing import Protocol -from app.domain.entities.signal import SignalMeta, SignalResult +from app.domain.entities.signal import SignalHit, SignalMeta, SignalResult class SignalRepository(Protocol): @@ -17,3 +17,6 @@ class SignalRepository(Protocol): def list_recent( self, as_of: date | None = None, limit: int = 20 ) -> list[SignalMeta]: ... + + def list_by_symbol(self, symbol: str, limit: int = 50) -> list[SignalHit]: + """该股在历史信号中的命中(Chart 标记用,含 signal_id 溯源)。""" diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py index 3f23265..13931a3 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/selection_impl.py @@ -12,6 +12,7 @@ from datetime import date, datetime from sqlalchemy import select from sqlalchemy.orm import Session +from app.domain.entities.chart import SelectionHit from app.domain.entities.selection import ( SelectionCandidate, SelectionMeta, @@ -84,6 +85,27 @@ class SqlAlchemySelectionRepository: 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( self, as_of: date | None = None, diff --git a/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py b/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py index 2c31f6b..0a04d6f 100644 --- a/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py +++ b/backend/app/infrastructure/persistence/sqlalchemy/repositories/signal_impl.py @@ -10,6 +10,7 @@ from sqlalchemy.orm import Session from app.domain.entities.signal import ( SignalEvent, + SignalHit, SignalMeta, SignalResult, SignalRules, @@ -87,6 +88,25 @@ class SqlAlchemySignalRepository: 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]: stmt = select(SignalSnapshotModel).order_by(SignalSnapshotModel.created_at.desc()) if as_of is not None: diff --git a/backend/tests/test_charts.py b/backend/tests/test_charts.py new file mode 100644 index 0000000..db293fc --- /dev/null +++ b/backend/tests/test_charts.py @@ -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