import sys import json import hashlib import tempfile import types import unittest from pathlib import Path from unittest import mock ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) sys.path.insert(0, str(ROOT / "src")) from platforms import qmt_research_runner class FakeFrame: def __len__(self): return 1 class QmtResearchRunnerTest(unittest.TestCase): def test_build_params_normalizes_official_timestamps(self): params = qmt_research_runner.build_qmt_parameters( stock_code="600000.sh", start_time="20260102", end_time="2026-01-30", ) self.assertEqual(params["stock_code"], "600000.SH") self.assertEqual(params["start_time"], "2026-01-02 00:00:00") self.assertEqual(params["end_time"], "2026-01-30 23:59:59") self.assertEqual(params["trade_mode"], "backtest") self.assertEqual(params["quote_mode"], "history") self.assertEqual(params["account_id"], "test") self.assertEqual(params["benchmark"], "000905.SH") baseline_bytes = (ROOT / "configs/baseline.json").read_bytes() baseline = json.loads(baseline_bytes) self.assertEqual( params["max_vol_rate"], baseline["execution"]["participation_rate"], ) self.assertEqual( params["slippage"], baseline["execution"]["slippage_bps"] / 10_000, ) expected_commission = ( baseline["fees"]["commission_rate"] + baseline["fees"]["transfer_fee_rate"] ) self.assertEqual(params["open_commission"], expected_commission) self.assertEqual(params["close_commission"], expected_commission) self.assertEqual( params["q60_baseline_config_sha256"], hashlib.sha256(baseline_bytes).hexdigest(), ) def test_live_mode_overrides_are_rejected(self): with self.assertRaises(ValueError): qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", overrides={"trade_mode": "trading"}, ) with self.assertRaises(ValueError): qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", overrides={"period": "1m"}, ) with self.assertRaises(ValueError): qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", overrides={"benchmark": "000300.SH"}, ) def test_non_finite_asset_and_rate_are_rejected(self): for kwargs in ( {"asset": float("nan")}, {"asset": float("inf")}, {"overrides": {"max_vol_rate": float("nan")}}, {"overrides": {"open_commission": float("inf")}}, ): with self.subTest(kwargs=kwargs): with self.assertRaisesRegex(ValueError, "finite"): qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", **kwargs, ) def test_intraday_period_is_rejected_to_prevent_daily_bar_lookahead(self): with self.assertRaisesRegex(ValueError, "period must be 1d"): qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", period="1m", ) def test_preflight_and_run_use_param_keyword(self): params = qmt_research_runner.build_qmt_parameters( stock_code="600000.SH", start_time="20260102", end_time="20260130", ) calls = {} def run_strategy_file(user_script, *, param): calls["user_script"] = user_script calls["param"] = param return object() fake_qmttools = types.SimpleNamespace(run_strategy_file=run_strategy_file) def get_market_data_ex(**kwargs): calls["data"] = kwargs return {"600000.SH": FakeFrame()} fake_xtdata = types.SimpleNamespace(get_market_data_ex=get_market_data_ex) fake_xtquant = types.SimpleNamespace(__version__="fake") def import_module(name): return { "xtquant": fake_xtquant, "xtquant.qmttools": fake_qmttools, "xtquant.xtdata": fake_xtdata, }[name] with tempfile.TemporaryDirectory() as temp: strategy = Path(temp) / "strategy.py" strategy.write_text("def init(C): pass\n", encoding="ascii") with mock.patch.object( qmt_research_runner.importlib.util, "find_spec", return_value=object(), ), mock.patch.object( qmt_research_runner.importlib, "import_module", side_effect=import_module, ): result = qmt_research_runner.run_qmt_backtest(strategy, params) self.assertIsNotNone(result) self.assertEqual(calls["param"], params) self.assertEqual(calls["data"]["start_time"], "20260102000000") self.assertEqual(calls["data"]["end_time"], "20260130235959") if __name__ == "__main__": unittest.main()