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