Files
qlib/backend/tests/test_charts.py
T
Simon 5bde8f9f5f feat(chart): M9-5 复权口径坐标换算落地(qfq 基准=最新因子 + marker 贴图换算)
- ChartService qfq 基准改为该股最新因子(截至今天)归一:历史区间随最新除权
  平移正确(v3 §20.5 Chart Display vs Execution basis 分离)
- 显示口径与执行 basis 不一致时,成交/信号 marker 价格按当日因子换算到 K 线坐标系
  (fill 早段价格在 qfq 下折算验证 100→50)
- tests:qfq 回测 marker 折算 + selection 标记保留;全量 pytest 通过
2026-09-09 07:15:41 +08:00

300 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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)