Files
quant-os/tests/platform_qmt_runner_test.py
T

156 lines
5.5 KiB
Python

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()