feat: operationalize local Tushare Qlib research
This commit is contained in:
@@ -0,0 +1,379 @@
|
||||
import hashlib
|
||||
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 tools.tushare_snapshot import (
|
||||
TushareSnapshotError,
|
||||
create_snapshot,
|
||||
write_manifest,
|
||||
)
|
||||
|
||||
|
||||
def _manifest_hash(manifest):
|
||||
payload = dict(manifest)
|
||||
payload.pop("manifest_sha256")
|
||||
encoded = json.dumps(
|
||||
payload,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
allow_nan=False,
|
||||
).encode()
|
||||
return hashlib.sha256(encoded).hexdigest()
|
||||
|
||||
|
||||
class TushareSnapshotTest(unittest.TestCase):
|
||||
def _new_mirror(self, root: Path):
|
||||
mirror = root / "mirror"
|
||||
(mirror / "data/parquet").mkdir(parents=True)
|
||||
connection = sqlite3.connect(mirror / "data/state.sqlite3")
|
||||
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,
|
||||
completed_at TEXT,
|
||||
params_json TEXT,
|
||||
fields TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
return mirror, connection
|
||||
|
||||
def _completed(
|
||||
self,
|
||||
mirror: Path,
|
||||
connection: sqlite3.Connection,
|
||||
*,
|
||||
job_id: int,
|
||||
api: str,
|
||||
partition: str,
|
||||
content: bytes = b"not-a-parquet",
|
||||
path_suffix: str = "",
|
||||
row_count: int = 1,
|
||||
):
|
||||
safe_partition = (
|
||||
partition.replace("/", "_").replace("=", "-") + path_suffix
|
||||
)
|
||||
path = mirror / "data/parquet" / api / f"{safe_partition}.parquet"
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(content)
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO jobs
|
||||
(id, api_name, partition_key, status, row_count, byte_count,
|
||||
sha256, file_path, completed_at, params_json, fields)
|
||||
VALUES (?, ?, ?, 'completed', ?, ?, ?, ?, ?, '[]', NULL)
|
||||
""",
|
||||
(
|
||||
job_id,
|
||||
api,
|
||||
partition,
|
||||
row_count,
|
||||
path.stat().st_size,
|
||||
hashlib.sha256(content).hexdigest(),
|
||||
str(path),
|
||||
f"2018-03-{min(job_id, 28):02d}T00:00:00Z",
|
||||
),
|
||||
)
|
||||
return path
|
||||
|
||||
def _core_fixture(self, root: Path):
|
||||
mirror, connection = self._new_mirror(root)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=1,
|
||||
api="daily",
|
||||
partition="month=2018-01",
|
||||
path_suffix="-old",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=2,
|
||||
api="daily",
|
||||
partition="month=2018-01",
|
||||
content=b"newer",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=3,
|
||||
api="daily",
|
||||
partition="month=2018-02",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=4,
|
||||
api="daily",
|
||||
partition="month=2018-03",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=5,
|
||||
api="adj_factor",
|
||||
partition="month=2018-01",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=6,
|
||||
api="adj_factor",
|
||||
partition="month=2018-02",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=7,
|
||||
api="trade_cal",
|
||||
partition="exchange=SSE/year=2018",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=8,
|
||||
api="stock_basic",
|
||||
partition="list_status=L",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=9,
|
||||
api="index_daily",
|
||||
partition="ts_code=000905_SH/year=2018",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=10,
|
||||
api="index_weight",
|
||||
partition="index_code=000905_SH/year=2018",
|
||||
)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=20,
|
||||
api="index_weight",
|
||||
partition="index_code=000905_SH/year=2017",
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO jobs
|
||||
(id, api_name, partition_key, status, row_count, byte_count,
|
||||
sha256, file_path, completed_at)
|
||||
VALUES
|
||||
(11, 'index_daily', 'ts_code=931643_CNI/year=2018', 'running',
|
||||
NULL, NULL, NULL, NULL, NULL)
|
||||
"""
|
||||
)
|
||||
connection.commit()
|
||||
return mirror, connection
|
||||
|
||||
def test_scoped_selection_shadows_old_completed_and_excludes_other_index(self):
|
||||
with tempfile.TemporaryDirectory() as temporary:
|
||||
root = Path(temporary)
|
||||
mirror, connection = self._core_fixture(root)
|
||||
connection.close()
|
||||
|
||||
manifest = create_snapshot(
|
||||
mirror,
|
||||
start="2018-01-15",
|
||||
end="2018-02-03",
|
||||
verify=False,
|
||||
)
|
||||
|
||||
self.assertEqual(manifest["summary"]["selected_jobs"], 9)
|
||||
self.assertEqual(
|
||||
manifest["summary"]["shadowed_completed_jobs"],
|
||||
1,
|
||||
)
|
||||
selected_ids = {job["id"] for job in manifest["selected_jobs"]}
|
||||
self.assertIn(2, selected_ids)
|
||||
self.assertNotIn(1, selected_ids)
|
||||
self.assertNotIn(4, selected_ids)
|
||||
self.assertIn(20, selected_ids)
|
||||
self.assertEqual(
|
||||
manifest["scope"]["index_weight_anchor_year"],
|
||||
2017,
|
||||
)
|
||||
self.assertEqual(
|
||||
manifest["excluded_live_jobs"][0]["reason"],
|
||||
"different_index_code",
|
||||
)
|
||||
self.assertEqual(
|
||||
manifest["verification"]["mode"],
|
||||
"not_performed",
|
||||
)
|
||||
self.assertTrue(
|
||||
manifest["verification"][
|
||||
"selected_jobs_still_latest_after_verification"
|
||||
]
|
||||
)
|
||||
self.assertRegex(
|
||||
manifest["selected_jobs"][0]["request_sha256"],
|
||||
r"^[0-9a-f]{64}$",
|
||||
)
|
||||
self.assertNotIn(
|
||||
"params_json",
|
||||
manifest["selected_jobs"][0],
|
||||
)
|
||||
self.assertEqual(
|
||||
manifest["manifest_sha256"],
|
||||
_manifest_hash(manifest),
|
||||
)
|
||||
serialized = json.dumps(manifest)
|
||||
self.assertNotIn(str(root), serialized)
|
||||
self.assertTrue(
|
||||
all(
|
||||
not Path(job["path"]).is_absolute()
|
||||
for job in manifest["selected_jobs"]
|
||||
)
|
||||
)
|
||||
|
||||
output = root / "snapshot.json"
|
||||
self.assertEqual(write_manifest(output, manifest), output.resolve())
|
||||
self.assertEqual(
|
||||
json.loads(output.read_text())["manifest_sha256"],
|
||||
manifest["manifest_sha256"],
|
||||
)
|
||||
|
||||
def test_in_scope_running_job_fails_closed(self):
|
||||
with tempfile.TemporaryDirectory() as temporary:
|
||||
root = Path(temporary)
|
||||
mirror, connection = self._core_fixture(root)
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO jobs
|
||||
(id, api_name, partition_key, status, row_count, byte_count,
|
||||
sha256, file_path, completed_at)
|
||||
VALUES
|
||||
(12, 'daily', 'month=2018-02', 'running',
|
||||
NULL, NULL, NULL, NULL, NULL)
|
||||
"""
|
||||
)
|
||||
connection.commit()
|
||||
connection.close()
|
||||
with self.assertRaisesRegex(
|
||||
TushareSnapshotError,
|
||||
"contains running/deferred",
|
||||
):
|
||||
create_snapshot(
|
||||
mirror,
|
||||
start="2018-01-15",
|
||||
end="2018-02-03",
|
||||
verify=False,
|
||||
)
|
||||
|
||||
def test_missing_partition_fails_closed(self):
|
||||
with tempfile.TemporaryDirectory() as temporary:
|
||||
root = Path(temporary)
|
||||
mirror, connection = self._new_mirror(root)
|
||||
self._completed(
|
||||
mirror,
|
||||
connection,
|
||||
job_id=1,
|
||||
api="daily",
|
||||
partition="month=2018-01",
|
||||
)
|
||||
connection.commit()
|
||||
connection.close()
|
||||
with self.assertRaisesRegex(
|
||||
TushareSnapshotError,
|
||||
"missing daily:month=2018-02",
|
||||
):
|
||||
create_snapshot(
|
||||
mirror,
|
||||
start="2018-01-01",
|
||||
end="2018-02-28",
|
||||
apis=("daily",),
|
||||
verify=False,
|
||||
)
|
||||
|
||||
@unittest.skipUnless(
|
||||
importlib.util.find_spec("pyarrow") is not None,
|
||||
"pyarrow is needed for Parquet metadata verification",
|
||||
)
|
||||
def test_default_verification_detects_file_tampering(self):
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
with tempfile.TemporaryDirectory() as temporary:
|
||||
root = Path(temporary)
|
||||
mirror, connection = self._new_mirror(root)
|
||||
table = pa.table(
|
||||
{
|
||||
"ts_code": ["600000.SH", "000001.SZ"],
|
||||
"trade_date": ["20180102", "20180102"],
|
||||
}
|
||||
)
|
||||
parquet_path = (
|
||||
mirror / "data/parquet/daily/month-2018-01.parquet"
|
||||
)
|
||||
parquet_path.parent.mkdir(parents=True)
|
||||
pq.write_table(table, parquet_path)
|
||||
content = parquet_path.read_bytes()
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO jobs
|
||||
(id, api_name, partition_key, status, row_count, byte_count,
|
||||
sha256, file_path, completed_at, params_json, fields)
|
||||
VALUES
|
||||
(1, 'daily', 'month=2018-01', 'completed', 2, ?, ?, ?,
|
||||
'2018-02-01T00:00:00Z', '[]',
|
||||
'ts_code,trade_date')
|
||||
""",
|
||||
(
|
||||
len(content),
|
||||
hashlib.sha256(content).hexdigest(),
|
||||
str(parquet_path),
|
||||
),
|
||||
)
|
||||
connection.commit()
|
||||
connection.close()
|
||||
|
||||
manifest = create_snapshot(
|
||||
mirror,
|
||||
start="2018-01-01",
|
||||
end="2018-01-31",
|
||||
apis=("daily",),
|
||||
)
|
||||
self.assertEqual(manifest["verification"]["verified_jobs"], 1)
|
||||
self.assertEqual(
|
||||
manifest["verification"]["double_verified_jobs"],
|
||||
1,
|
||||
)
|
||||
self.assertEqual(
|
||||
manifest["selected_jobs"][0]["observed"]["row_count"],
|
||||
2,
|
||||
)
|
||||
|
||||
parquet_path.write_bytes(content + b"tampered")
|
||||
with self.assertRaisesRegex(TushareSnapshotError, "size mismatch"):
|
||||
create_snapshot(
|
||||
mirror,
|
||||
start="2018-01-01",
|
||||
end="2018-01-31",
|
||||
apis=("daily",),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user