feat: add Tushare Qlib data bridge

This commit is contained in:
2026-07-26 13:41:24 +08:00
parent 48c5f64bbd
commit e5911ca454
15 changed files with 1804 additions and 8 deletions
+3
View File
@@ -119,11 +119,14 @@ class CapabilityProbeTest(unittest.TestCase):
self.assertTrue(report["artifacts"]["portable_core"])
self.assertTrue(report["artifacts"]["colab_notebook"])
self.assertTrue(report["artifacts"]["jqdata_snapshot_adapter"])
self.assertTrue(report["artifacts"]["tushare_qlib_adapter"])
self.assertTrue(report["artifacts"]["tushare_qlib_cli"])
self.assertTrue(report["artifacts"]["qmt_shadow_planner"])
self.assertTrue(report["artifacts"]["qmt_shadow_cli"])
self.assertIn("no external platform", report["verification_boundary"])
self.assertTrue(report["matrix"]["xttrader"]["safe_default"])
self.assertIn("rejects BJ", report["matrix"]["xttrader"]["market_scope"])
self.assertTrue(report["matrix"]["tushare_local_qlib"]["safe_default"])
for probe in report["optional_runtimes"].values():
self.assertIs(probe["import_ok"], None)
+35
View File
@@ -17,6 +17,7 @@ sys.path.insert(0, str(ROOT / "src"))
from platforms.qlib_runner import (
QlibUnavailableError,
_main,
_validate_managed_provider_request,
build_alpha158_lightgbm_task,
generate_portable_momentum_signal,
run_native_momentum_backtest,
@@ -91,6 +92,35 @@ class QlibFixtureTest(unittest.TestCase):
end_time="2024-02-01",
)
def test_managed_provider_request_fails_before_terminal_session(self):
evidence = {
"managed": True,
"calendar": {
"start": "2024-01-02",
"end": "2024-01-11",
},
"market": {"name": "tushare_a"},
"benchmark": {"symbol": "SH999999"},
}
with self.assertRaisesRegex(
ValueError,
"simulator consumes the next provider session",
):
_validate_managed_provider_request(
evidence,
market="tushare_a",
benchmark="SH999999",
start_time="2024-01-02",
end_time="2024-01-11",
)
_validate_managed_provider_request(
evidence,
market="tushare_a",
benchmark="SH999999",
start_time="2024-01-02",
end_time="2024-01-10",
)
@unittest.skipUnless(
importlib.util.find_spec("pandas") is not None,
"pandas is optional outside the Qlib environment",
@@ -147,6 +177,7 @@ class QlibCliTest(unittest.TestCase):
"clock_contract": "test-clock",
"strategy_fidelity": "signal-only",
"a_share_rule_fidelity": "approximation",
"run_parameters": {"start_time": "2024-01-02"},
}
with tempfile.TemporaryDirectory() as temp:
output = Path(temp) / "nested" / "summary.json"
@@ -194,6 +225,10 @@ class QlibCliTest(unittest.TestCase):
self.assertEqual(payload["workflow"], "momentum")
self.assertEqual(payload["signal_rows"], 2)
self.assertEqual(payload["portfolio_frequencies"], ["1day"])
self.assertEqual(
payload["run_parameters"]["start_time"],
"2024-01-02",
)
self.assertIn('"workflow": "momentum"', printer.call_args.args[0])
def test_alpha158_cli_dispatches_with_validated_segments(self):
+230
View File
@@ -0,0 +1,230 @@
import importlib.util
import json
from pathlib import Path
import sqlite3
import sys
import tempfile
import unittest
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from adapters.tushare_local import (
TushareMirrorError,
build_qlib_provider,
tushare_to_qlib_symbol,
verify_qlib_provider,
)
class TushareSymbolTest(unittest.TestCase):
def test_symbol_conversion(self):
self.assertEqual(tushare_to_qlib_symbol("600000.SH"), "SH600000")
self.assertEqual(tushare_to_qlib_symbol("000001.SZ"), "SZ000001")
self.assertEqual(tushare_to_qlib_symbol("430047.BJ"), "BJ430047")
with self.assertRaises(ValueError):
tushare_to_qlib_symbol("T600018.SH")
@unittest.skipUnless(
importlib.util.find_spec("pandas") is not None
and importlib.util.find_spec("pyarrow") is not None,
"pandas and pyarrow are optional outside the Qlib research environment",
)
class TushareProviderTest(unittest.TestCase):
def _mirror(self, root: Path, *, bad_ohlc: bool = False) -> Path:
import pandas
mirror = root / "mirror"
parquet_root = mirror / "data/parquet"
(parquet_root / "daily").mkdir(parents=True)
(parquet_root / "trade_cal").mkdir(parents=True)
(parquet_root / "stock_basic").mkdir(parents=True)
dates = pandas.bdate_range("2024-01-02", periods=8)
daily_rows = []
for symbol, base in (("600000.SH", 10.0), ("000001.SZ", 20.0)):
previous = base
for index, value in enumerate(dates):
close = base * (1.01 ** index)
daily_rows.append(
{
"ts_code": symbol,
"trade_date": value.strftime("%Y%m%d"),
"open": close * 0.99,
"high": close * 1.01,
"low": close * 0.98,
"close": close,
"pre_close": previous,
"change": close - previous,
"pct_chg": (close / previous - 1.0) * 100.0,
"vol": 1000.0,
"amount": close * 100.0,
}
)
previous = close
if bad_ohlc:
daily_rows[0]["high"] = 8.0
daily_rows[0]["low"] = 9.0
frames = {
"daily/month=2024-01.parquet": pandas.DataFrame(daily_rows),
"trade_cal/exchange=SSE_year=2024.parquet": pandas.DataFrame(
{
"exchange": ["SSE"] * len(dates),
"cal_date": [value.strftime("%Y%m%d") for value in dates],
"is_open": [1] * len(dates),
"pretrade_date": [None] * len(dates),
}
),
"stock_basic/list_status=L.parquet": pandas.DataFrame(
{
"ts_code": ["600000.SH", "000001.SZ"],
"curr_type": ["CNY", "CNY"],
}
),
}
state = mirror / "data/state.sqlite3"
connection = sqlite3.connect(state)
connection.execute(
"""
CREATE TABLE jobs (
id INTEGER PRIMARY KEY,
api_name TEXT NOT NULL,
partition_key TEXT NOT NULL,
status TEXT NOT NULL,
row_count INTEGER,
byte_count INTEGER,
sha256 TEXT,
file_path TEXT,
started_at TEXT,
completed_at TEXT,
error_code INTEGER
)
"""
)
import hashlib
for index, (relative, frame) in enumerate(frames.items(), start=1):
path = parquet_root / relative
path.parent.mkdir(parents=True, exist_ok=True)
frame.to_parquet(path, index=False)
api = relative.split("/", 1)[0]
connection.execute(
"""
INSERT INTO jobs VALUES (?, ?, ?, 'completed', ?, ?, ?, ?,
NULL, '2024-01-31T00:00:00Z', NULL)
""",
(
index,
api,
path.stem,
len(frame),
path.stat().st_size,
hashlib.sha256(path.read_bytes()).hexdigest(),
str(path),
),
)
connection.execute(
"""
INSERT INTO jobs VALUES
(99, 'daily', 'month=2024-02', 'running', NULL, NULL, NULL, NULL,
'2024-02-01T00:00:00Z', NULL, NULL)
"""
)
connection.commit()
connection.close()
return mirror
def test_build_and_verify_ignores_running_job(self):
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
mirror = self._mirror(root)
output = root / "provider"
result = build_qlib_provider(
mirror,
output,
start="2024-01-02",
end="2024-01-11",
minimum_observations=2,
allow_unadjusted=True,
)
self.assertFalse(result["investment_value_claim"])
self.assertEqual(result["market"]["instrument_count"], 2)
self.assertTrue(verify_qlib_provider(output)["ok"])
manifest = json.loads(
(output / "quant_os_tushare_manifest.json").read_text()
)
self.assertRegex(
manifest["converter"]["source_sha256"],
r"^[0-9a-f]{64}$",
)
self.assertEqual(
manifest["mirror_job_status_counts_at_build"]["running"],
1,
)
manifest["build_parameters"]["start"] = "2024-01-03"
(output / "quant_os_tushare_manifest.json").write_text(
json.dumps(manifest),
encoding="utf-8",
)
with self.assertRaisesRegex(
TushareMirrorError,
"data_version",
):
verify_qlib_provider(output)
def test_unadjusted_and_ohlc_gates_fail_closed(self):
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
mirror = self._mirror(root)
with self.assertRaises(TushareMirrorError):
build_qlib_provider(
mirror,
root / "provider",
start="2024-01-02",
end="2024-01-11",
minimum_observations=2,
)
def test_expand_range_uses_the_original_ohlc_envelope(self):
from array import array
with tempfile.TemporaryDirectory() as temporary:
root = Path(temporary)
mirror = self._mirror(root, bad_ohlc=True)
with self.assertRaisesRegex(
TushareMirrorError,
"OHLC envelope anomalies",
):
build_qlib_provider(
mirror,
root / "provider-fail",
start="2024-01-02",
end="2024-01-11",
minimum_observations=2,
allow_unadjusted=True,
)
output = root / "provider-expanded"
build_qlib_provider(
mirror,
output,
start="2024-01-02",
end="2024-01-11",
minimum_observations=2,
allow_unadjusted=True,
expand_range=True,
)
high = array("f")
high.frombytes(
(output / "features/sh600000/high.day.bin").read_bytes()
)
low = array("f")
low.frombytes(
(output / "features/sh600000/low.day.bin").read_bytes()
)
self.assertAlmostEqual(high[1], 10.0)
self.assertAlmostEqual(low[1], 8.0)
if __name__ == "__main__":
unittest.main()