Files
qlib/backend/tests/test_charts.py
T
Simon 995ed08548 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 通过
2026-09-09 07:09:52 +08:00

234 lines
9.5 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