"""股票名称(`name`)填充与 `GET /api/stocks/names` 接口测试。 覆盖三件事(需求:前端任何出现代码的地方都要能显示名称): 1) 回测结果的展示结构(symbol_curves / positions / trades / fills / signal_history / selection_history)由 `ResearchService.run_backtest` 统一回填名称; 股票池为空时**静默跳过**(名称只是展示增强,不得让已算完的回测失败); 2) 选股结果 `candidates[].name` 由 `SelectionService` 用已装配股票池回填; 3) `GET /api/stocks/names` 返回 `{symbol: name}` **dict**(前端契约), 且 `GET /api/stocks/{symbol}` 未被 `/names` 破坏(路由顺序)。 用内存 Repository + 合成行情(不连真库、不跑长区间),秒级完成。 """ from __future__ import annotations from datetime import date import pandas as pd import pytest from app.api import deps from app.application.services.selection_service import SelectionService from app.domain.entities.market import Stock from app.domain.entities.research import ( FactorSpec, ResearchSpec, SelectionSpec, UniverseSpec, ) from app.domain.entities.selection import SelectionQuery from app.main import app from app.quant.engine import LocalEngine from app.quant.service import ResearchService, _fill_names from fastapi.testclient import TestClient from conftest_quant import bars_dataframe_to_daily_bars, synthetic_daily _SYMS = ["600000.SH", "600001.SH", "600002.SH"] _NAMES = {"600000.SH": "浦发银行", "600001.SH": "邯郸钢铁", "600002.SH": "齐鲁石化"} def _stocks(symbols: list[str] | None = None) -> list[Stock]: return [ Stock(symbol=s, name=_NAMES[s], list_date=date(1999, 11, 10)) for s in (symbols or _SYMS) ] class _MemStockRepo: def __init__(self, stocks: list[Stock]) -> None: self._stocks = stocks def list(self) -> list[Stock]: return self._stocks def get_by_symbol(self, symbol: str) -> Stock | None: return next((s for s in self._stocks if s.symbol == symbol), None) class _MemDailyRepo: def __init__(self, df: pd.DataFrame) -> None: self._bars = bars_dataframe_to_daily_bars(df) def get_range(self, symbol, start, end): return [b for b in self._bars if b.symbol == symbol and start <= b.trade_date <= end] def get_range_many(self, symbols, start, end, adjust="none"): syms = set(symbols) return [b for b in self._bars if b.symbol in syms and start <= b.trade_date <= end] def _daily_df() -> pd.DataFrame: return synthetic_daily( {"600000.SH": 0.004, "600001.SH": 0.001, "600002.SH": -0.003}, n=300 ) def _backtest_spec() -> ResearchSpec: return ResearchSpec( type="backtest", universe=UniverseSpec(exclude_st=False, min_listing_days=0), factors=[FactorSpec(name="momentum_20", weight=1.0)], selection=SelectionSpec(top_n=2), rebalance="monthly", # 半年区间(约 150 个交易日):够触发月度调仓与卖出,又足够快 period=(date(2024, 6, 3), date(2024, 12, 31)), ) class TestBacktestNameFill: def test_all_display_structures_carry_name(self) -> None: service = ResearchService( _MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine() ) result = service.run_backtest(_backtest_spec()) assert result.symbol_curves, "应持有过股票并输出个股曲线" assert result.positions, "应有调仓后仓位记录" assert result.selection_history, "应有择股记录" assert result.signal_history, "应有交易意图记录" for curve in result.symbol_curves: assert curve.name == _NAMES[curve.symbol] for pos in result.positions: assert pos.name == _NAMES[pos.symbol] for pick in result.selection_history: assert pick.name == _NAMES[pick.symbol] # 成交/意图记录里存在 symbol="" 的池子不足提示(非股票),其 name 保持 None for action in result.signal_history + result.fills: if action.symbol: assert action.name == _NAMES[action.symbol] else: assert action.name is None def test_trades_and_marks_carry_name(self) -> None: service = ResearchService( _MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine() ) result = service.run_backtest(_backtest_spec()) assert result.trades, "月度调仓 + 因子换手应产生已卖出往返" for trade in result.trades: assert trade.name == _NAMES[trade.symbol] # marks 与 signal_history 同源,也应带上名称(前端个股曲线标注用) for curve in result.symbol_curves: for mark in curve.marks: assert mark.name == _NAMES[mark.symbol] def test_fill_names_skips_silently_when_stock_pool_empty(self) -> None: """空股票池 → 静默跳过(名称缺失只影响展示,不得抛错)。""" service = ResearchService( _MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine() ) result = service.run_backtest(_backtest_spec()) before = result.symbol_curves[0].name returned = _fill_names(result, []) assert returned is result assert result.symbol_curves[0].name == before # 未被清空、也未被改写 def test_fill_names_is_idempotent_and_keeps_name_without_stock_record(self) -> None: """名称回填幂等;股票池里没有该代码时保持 None(不伪造名称)。""" service = ResearchService( _MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), LocalEngine() ) result = service.run_backtest(_backtest_spec()) snapshot = [(c.symbol, c.name) for c in result.symbol_curves] _fill_names(result, _stocks()) assert [(c.symbol, c.name) for c in result.symbol_curves] == snapshot only_other = [Stock(symbol="601398.SH", name="工商银行", list_date=date(2006, 1, 1))] _fill_names(result, only_other) assert all(c.name == _NAMES[c.symbol] for c in result.symbol_curves) class TestSelectionCandidateName: def _service(self) -> SelectionService: return SelectionService(_MemStockRepo(_stocks()), _MemDailyRepo(_daily_df())) def test_score_mode_candidates_carry_name(self) -> None: result = self._service().select( SelectionQuery( method="score", factors=[{"name": "momentum_20", "weight": 1.0}], top_n=2, as_of=date(2024, 12, 31), universe=UniverseSpec(exclude_st=False, min_listing_days=0), ) ) assert result.candidates for cand in result.candidates: assert cand.name == _NAMES[cand.symbol] def test_condition_mode_candidates_carry_name(self) -> None: result = self._service().select( SelectionQuery( method="condition", conditions=[{"field": "close", "op": "gte", "ref": "ma20"}], as_of=date(2024, 12, 31), universe=UniverseSpec(exclude_st=False, min_listing_days=0), ) ) assert result.candidates for cand in result.candidates: assert cand.name == _NAMES[cand.symbol] class TestQlibEngineNameFill: def test_qlib_engine_path_also_fills_names(self, tmp_path) -> None: """名称回填写在服务层,因此 LocalEngine / QlibEngine 两条路径都必须生效。 这条用例走真实的 Qlib 数据管线(落盘 bin → D.features 读回), 而不是只测本地引擎后「假设」另一条路径也覆盖到了。 """ pytest.importorskip("qlib") # 未安装 qlib 的环境跳过,不假装覆盖 from app.quant.qlib_adapter.engine import QlibEngine service = ResearchService( _MemStockRepo(_stocks()), _MemDailyRepo(_daily_df()), QlibEngine(qlib_dir=tmp_path) ) result = service.run_backtest(_backtest_spec()) assert result.symbol_curves, "qlib 路径也应产生个股曲线" for curve in result.symbol_curves: assert curve.name == _NAMES[curve.symbol] for pos in result.positions: assert pos.name == _NAMES[pos.symbol] for pick in result.selection_history: assert pick.name == _NAMES[pick.symbol] @pytest.fixture() def client() -> TestClient: """只覆盖 stock 仓储(这三个接口只用 StockRepoDep,不需要 DB session)。""" stocks = [*_stocks(), Stock(symbol="600519.SH", name="贵州茅台", list_date=date(2001, 8, 27))] app.dependency_overrides[deps._stock_repo_factory] = lambda: _MemStockRepo(stocks) # noqa: SLF001 with TestClient(app) as c: yield c app.dependency_overrides.clear() class TestStockNamesApi: def test_names_returns_symbol_to_name_dict(self, client: TestClient) -> None: resp = client.get("/api/stocks/names") assert resp.status_code == 200 body = resp.json() # 契约:dict(不是数组),前端按 Record 消费 assert isinstance(body, dict) assert body["600519.SH"] == "贵州茅台" assert body["600000.SH"] == _NAMES["600000.SH"] assert set(body) == {"600000.SH", "600001.SH", "600002.SH", "600519.SH"} def test_symbol_path_not_swallowed_by_names_route(self, client: TestClient) -> None: """路由顺序回归:/{symbol} 仍按代码查询(证明 /names 未被参数吞掉、也没吞掉它)。""" one = client.get("/api/stocks/600519.SH") assert one.status_code == 200 assert one.json()["symbol"] == "600519.SH" assert one.json()["name"] == "贵州茅台" assert client.get("/api/stocks/999999.SZ").status_code == 404 assert client.get("/api/stocks?q=茅台").json()[0]["symbol"] == "600519.SH"