feat: add Quant OS A-share baseline
This commit is contained in:
@@ -0,0 +1,155 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user