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, inventory_mirror, tushare_to_qlib_symbol, verify_qlib_provider, ) from tools.tushare_qlib import _assert_json_outside_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") def test_provider_json_output_cannot_mutate_provider_tree(self): with self.assertRaisesRegex( ValueError, "outside the immutable provider", ): _assert_json_outside_provider( "/tmp/provider/evidence.json", "/tmp/provider", ) _assert_json_outside_provider( "/tmp/evidence.json", "/tmp/provider", ) @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, advanced: bool = False, shadowed_daily: 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 = [] symbols = [("600000.SH", 10.0), ("000001.SZ", 20.0)] if advanced: symbols.append(("000002.SZ", 30.0)) for symbol, base in symbols: 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": [value[0] for value in symbols], "curr_type": ["CNY"] * len(symbols), } ), } partitions = {} if advanced: adj_rows = [] for symbol, _ in symbols: for index, value in enumerate(dates): factor = 2.0 if symbol == "600000.SH" and index >= 4 else 1.0 adj_rows.append( { "ts_code": symbol, "trade_date": value.strftime("%Y%m%d"), "adj_factor": factor, } ) index_rows = [] previous = 1000.0 for index, value in enumerate(dates): close = 1000.0 + index * 10.0 index_rows.append( { "ts_code": "000905.SH", "trade_date": value.strftime("%Y%m%d"), "open": close - 2.0, "high": close + 5.0, "low": close - 5.0, "close": close, "pre_close": previous, "change": close - previous, "pct_chg": (close / previous - 1.0) * 100.0, "vol": 2000.0, "amount": close * 200.0, } ) previous = close first_snapshot = dates[0].strftime("%Y%m%d") second_snapshot = dates[4].strftime("%Y%m%d") anchor_weight_rows = [ { "index_code": "000905.SH", "con_code": "600000.SH", "trade_date": "20231229", "weight": 50.0, }, { "index_code": "000905.SH", "con_code": "000001.SZ", "trade_date": "20231229", "weight": 50.0, }, ] weight_rows = [ { "index_code": "000905.SH", "con_code": "600000.SH", "trade_date": first_snapshot, "weight": 50.0, }, { "index_code": "000905.SH", "con_code": "000001.SZ", "trade_date": first_snapshot, "weight": 50.0, }, { "index_code": "000905.SH", "con_code": "000001.SZ", "trade_date": second_snapshot, "weight": 55.0, }, { "index_code": "000905.SH", "con_code": "000002.SZ", "trade_date": second_snapshot, "weight": 45.0, }, ] frames.update( { "adj_factor/month=2024-01.parquet": pandas.DataFrame( adj_rows ), ( "index_daily/" "ts_code=000905_SH_year=2024.parquet" ): pandas.DataFrame(index_rows), ( "index_weight/" "index_code=000905_SH_year=2023.parquet" ): pandas.DataFrame(anchor_weight_rows), ( "index_weight/" "index_code=000905_SH_year=2024.parquet" ): pandas.DataFrame(weight_rows), } ) partitions.update( { ( "index_daily/" "ts_code=000905_SH_year=2024.parquet" ): "ts_code=000905_SH/year=2024", ( "index_weight/" "index_code=000905_SH_year=2023.parquet" ): "index_code=000905_SH/year=2023", ( "index_weight/" "index_code=000905_SH_year=2024.parquet" ): "index_code=000905_SH/year=2024", } ) 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, partitions.get(relative, path.stem), len(frame), path.stat().st_size, hashlib.sha256(path.read_bytes()).hexdigest(), str(path), ), ) if shadowed_daily: daily_path = parquet_root / "daily/month=2024-01.parquet" connection.execute( """ INSERT INTO jobs VALUES (0, 'daily', 'month=2024-01', 'completed', 1, 1, 'obsolete-checksum', ?, NULL, '2024-01-01T00:00:00Z', NULL) """, (str(daily_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) verified = verify_qlib_provider(output) self.assertTrue(verified["ok"]) self.assertTrue(verified["converter_source_matches_current"]) 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) def test_inventory_ignores_shadowed_completed_file_version(self): with tempfile.TemporaryDirectory() as temporary: mirror = self._mirror( Path(temporary), shadowed_daily=True, ) result = inventory_mirror(mirror) self.assertEqual(result["shadowed_completed_versions"], 1) self.assertEqual( result["shadowed_completed_versions_by_api"], {"daily": 1}, ) self.assertEqual(result["tables"]["daily"]["files"], 1) self.assertEqual(result["verified_completed_files"], 3) def test_adjusted_real_benchmark_and_pit_universe(self): from array import array import math with tempfile.TemporaryDirectory() as temporary: root = Path(temporary) mirror = self._mirror(root, advanced=True) output = root / "provider" result = build_qlib_provider( mirror, output, start="2024-01-02", end="2024-01-11", minimum_observations=999, benchmark_index_code="000905.SH", universe_index_code="000905.SH", expected_constituent_count=2, ) self.assertEqual(result["schema_version"], 2) self.assertEqual( result["price_adjustment"]["mode"], "cumulative_adjusted_first_observation_anchor", ) self.assertTrue(result["price_adjustment"]["component_complete"]) self.assertFalse(result["price_adjustment"]["production_eligible"]) self.assertEqual( result["benchmark"]["kind"], "tushare_index_daily", ) self.assertTrue(result["market"]["point_in_time"]) self.assertFalse(result["market"]["minimum_observations_applied"]) self.assertFalse(result["market"]["same_session_membership_use"]) self.assertEqual( result["market"]["availability_contract"], "conservative next-session effectiveness because the source " "does not expose a separately verified published_at timestamp", ) self.assertEqual(result["market"]["instrument_count"], 3) self.assertTrue( result["source_stability"]["post_read_file_verification"] ) self.assertEqual( {item["api_name"] for item in result["source_jobs"]}, { "adj_factor", "daily", "index_daily", "index_weight", "stock_basic", "trade_cal", }, ) factor = array("f") factor.frombytes( (output / "features/sh600000/factor.day.bin").read_bytes() ) close = array("f") close.frombytes( (output / "features/sh600000/close.day.bin").read_bytes() ) volume = array("f") volume.frombytes( (output / "features/sh600000/volume.day.bin").read_bytes() ) change = array("f") change.frombytes( (output / "features/sh600000/change.day.bin").read_bytes() ) self.assertAlmostEqual(factor[1], 1.0) self.assertAlmostEqual(factor[5], 2.0) raw_close = 10.0 * (1.01 ** 4) self.assertAlmostEqual(close[5] / factor[5], raw_close, places=5) self.assertAlmostEqual(volume[5], 50_000.0, places=3) self.assertTrue(math.isnan(change[1])) self.assertAlmostEqual( change[5], (raw_close * 2.0) / (10.0 * (1.01 ** 3)) - 1.0, places=5, ) index_close = array("f") index_close.frombytes( (output / "features/sh000905/close.day.bin").read_bytes() ) self.assertAlmostEqual(index_close[1], 1000.0) market_lines = ( output / "instruments/tushare_a.txt" ).read_text(encoding="utf-8").splitlines() self.assertEqual( market_lines, [ "SZ000001\t2024-01-02\t2024-01-11", "SZ000002\t2024-01-09\t2024-01-11", "SH600000\t2024-01-02\t2024-01-08", ], ) self.assertTrue(verify_qlib_provider(output)["ok"]) def test_pit_universe_fails_closed_on_missing_anchor_or_bad_roster(self): with tempfile.TemporaryDirectory() as temporary: root = Path(temporary) mirror = self._mirror(root, advanced=True) state = mirror / "data/state.sqlite3" connection = sqlite3.connect(state) connection.execute( "DELETE FROM jobs WHERE partition_key = ?", ("index_code=000905_SH/year=2023",), ) connection.commit() connection.close() with self.assertRaisesRegex( TushareMirrorError, "missing required partitions", ): build_qlib_provider( mirror, root / "provider-missing-anchor", start="2024-01-02", end="2024-01-11", benchmark_index_code="000905.SH", universe_index_code="000905.SH", expected_constituent_count=2, ) with tempfile.TemporaryDirectory() as temporary: root = Path(temporary) mirror = self._mirror(root, advanced=True) with self.assertRaisesRegex( TushareMirrorError, "constituents; expected 3", ): build_qlib_provider( mirror, root / "provider-bad-roster", start="2024-01-02", end="2024-01-11", benchmark_index_code="000905.SH", universe_index_code="000905.SH", expected_constituent_count=3, ) with self.assertRaisesRegex( ValueError, "requires benchmark_index_code", ): build_qlib_provider( mirror, root / "provider-synthetic-benchmark", start="2024-01-02", end="2024-01-11", universe_index_code="000905.SH", expected_constituent_count=2, ) if __name__ == "__main__": unittest.main()