- ChartService qfq 基准改为该股最新因子(截至今天)归一:历史区间随最新除权 平移正确(v3 §20.5 Chart Display vs Execution basis 分离) - 显示口径与执行 basis 不一致时,成交/信号 marker 价格按当日因子换算到 K 线坐标系 (fill 早段价格在 qfq 下折算验证 100→50) - tests:qfq 回测 marker 折算 + selection 标记保留;全量 pytest 通过
300 lines
12 KiB
Python
300 lines
12 KiB
Python
"""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
|
||
|
||
|
||
class TestAdjustCoordinate:
|
||
def test_qfq_backtest_marker_converted(self, seeded, tmp_path) -> None:
|
||
"""显示前复权时,早期成交 marker 价格按因子折算贴图(v3 §20.5 坐标)。"""
|
||
|
||
from app.domain.entities.research import (
|
||
ActionRecord,
|
||
BacktestResult,
|
||
BacktestSummary,
|
||
CurvePoint,
|
||
RankedPick,
|
||
Trade,
|
||
)
|
||
|
||
engine, Session, df = _seeded_with_factors(tmp_path)
|
||
dates = sorted(df["trade_date"].unique())
|
||
split = dates[len(dates) // 2]
|
||
with Session() as session:
|
||
SqlAlchemyAdjustFactorRepository(session).upsert_many(
|
||
[
|
||
AdjustFactor(symbol=_SYMS[0], trade_date=d,
|
||
factor=Decimal("1.0") if d < split else Decimal("2.0"))
|
||
for d in dates
|
||
]
|
||
)
|
||
session.commit()
|
||
svc = ChartService(
|
||
SqlAlchemyStockRepository(session),
|
||
SqlAlchemyDailyBarRepository(session),
|
||
SqlAlchemyAdjustFactorRepository(session),
|
||
)
|
||
# 早期(factor=1.0 段)一笔买入成交价 100
|
||
result = BacktestResult(
|
||
summary=BacktestSummary(
|
||
start=dates[0], end=dates[-1], initial_capital=1e6, final_equity=1e6,
|
||
total_return_pct=0, annual_return_pct=0, sharpe=0, max_drawdown_pct=0,
|
||
volatility_pct=0, win_rate_pct=0, total_trades=1, avg_turnover_pct=0,
|
||
),
|
||
equity_curve=[CurvePoint(date=dates[0], value=1e6)],
|
||
drawdown=[], monthly_returns=[], yearly_returns=[],
|
||
positions=[], trades=[
|
||
Trade(entry_date=dates[5], exit_date=dates[-1], symbol=_SYMS[0],
|
||
entry_price=100.0, exit_price=200.0, return_pct=100),
|
||
],
|
||
selection_history=[
|
||
RankedPick(date=dates[5], symbol=_SYMS[0], rank=1, score=1.0)
|
||
],
|
||
signal_history=[
|
||
ActionRecord(date=dates[5], symbol=_SYMS[0], signal="BUY", filled=True, price=100.0),
|
||
],
|
||
fills=[
|
||
ActionRecord(date=dates[5], symbol=_SYMS[0], signal="BUY", filled=True, price=100.0),
|
||
],
|
||
turnover_pct=0,
|
||
config_snapshot={"price_adjustment": "none"},
|
||
)
|
||
chart_none = svc.backtest_stock_chart(result, _SYMS[0], dates[0], dates[-1], "none")
|
||
chart_qfq = svc.backtest_stock_chart(result, _SYMS[0], dates[0], dates[-1], "qfq")
|
||
# none:fill 价 100;qfq 显示:早期因子 1.0 / 最新 2.0 → 折算 50
|
||
none_fill = chart_none.fills[0]
|
||
qfq_fill = chart_qfq.fills[0]
|
||
assert none_fill.price == 100.0
|
||
assert qfq_fill.price is not None and abs(qfq_fill.price - 50.0) < 1e-3
|
||
# 选股意图标记保留
|
||
assert any(m.kind == "selection" for m in chart_qfq.selections)
|