feat: preserve Quant OS target-package vertical slice
This commit is contained in:
@@ -0,0 +1,191 @@
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from quant60.features import TrainOnlyFeaturePreprocessor
|
||||
from quant60.ledger import canonical_json
|
||||
from quant60.model_bundle import (
|
||||
ModelBundleIntegrityError,
|
||||
create_model_bundle,
|
||||
load_verified_model_bundle,
|
||||
model_bundle_sha256,
|
||||
verify_model_bundle,
|
||||
write_model_bundle,
|
||||
)
|
||||
from quant60.research import DeterministicRidge
|
||||
|
||||
|
||||
HASHES = {
|
||||
"data_sha256": hashlib.sha256(b"training-data").hexdigest(),
|
||||
"source_sha256": hashlib.sha256(b"source-tree").hexdigest(),
|
||||
"config_sha256": hashlib.sha256(b"training-config").hexdigest(),
|
||||
}
|
||||
|
||||
|
||||
class ModelBundleTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.training_rows = [
|
||||
{"momentum": -0.2, "volatility": 0.10},
|
||||
{"momentum": -0.1, "volatility": None},
|
||||
{"momentum": 0.0, "volatility": 0.20},
|
||||
{"momentum": 0.1, "volatility": 0.25},
|
||||
{"momentum": 0.2, "volatility": 0.30},
|
||||
]
|
||||
self.processor = TrainOnlyFeaturePreprocessor(
|
||||
winsor_fraction=0.1
|
||||
).fit(self.training_rows)
|
||||
transformed = self.processor.transform(self.training_rows)
|
||||
matrix = [
|
||||
[row[name] for name in self.processor.output_names]
|
||||
for row in transformed
|
||||
]
|
||||
self.model = DeterministicRidge(alpha=0.5).fit(
|
||||
matrix,
|
||||
[-0.03, -0.01, 0.0, 0.02, 0.04],
|
||||
)
|
||||
self.bundle = create_model_bundle(
|
||||
self.processor,
|
||||
self.model,
|
||||
trained_through="2024-12-31",
|
||||
label_name="t_plus_1_open_to_t_plus_6_open_excess",
|
||||
horizon_trading_days=5,
|
||||
**HASHES,
|
||||
)
|
||||
|
||||
def test_freezes_preprocessor_model_label_and_lineage(self):
|
||||
verified = verify_model_bundle(self.bundle)
|
||||
self.assertEqual(
|
||||
verified.feature_order,
|
||||
tuple(sorted(self.training_rows[0])),
|
||||
)
|
||||
self.assertEqual(
|
||||
verified.output_feature_order,
|
||||
self.processor.output_names,
|
||||
)
|
||||
self.assertEqual(
|
||||
verified.coefficients,
|
||||
self.model.coef_,
|
||||
)
|
||||
self.assertEqual(verified.intercept, self.model.intercept_)
|
||||
self.assertEqual(verified.trained_through.isoformat(), "2024-12-31")
|
||||
self.assertEqual(verified.horizon_trading_days, 5)
|
||||
self.assertEqual(verified.data_sha256, HASHES["data_sha256"])
|
||||
|
||||
def test_write_read_and_prediction_replay_are_deterministic(self):
|
||||
replay_rows = [
|
||||
{"volatility": None, "momentum": 0.07},
|
||||
{"momentum": -100.0, "volatility": 100.0},
|
||||
]
|
||||
transformed = self.processor.transform(replay_rows)
|
||||
matrix = [
|
||||
[row[name] for name in self.processor.output_names]
|
||||
for row in transformed
|
||||
]
|
||||
expected = self.model.predict(matrix)
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
left = write_model_bundle(
|
||||
self.bundle,
|
||||
Path(directory) / "left.json",
|
||||
)
|
||||
right = write_model_bundle(
|
||||
copy.deepcopy(self.bundle),
|
||||
Path(directory) / "right.json",
|
||||
)
|
||||
self.assertEqual(left.read_bytes(), right.read_bytes())
|
||||
loaded = load_verified_model_bundle(
|
||||
left,
|
||||
expected_bundle_sha256=self.bundle["bundle_sha256"],
|
||||
**{
|
||||
f"expected_{key}": value
|
||||
for key, value in HASHES.items()
|
||||
},
|
||||
)
|
||||
self.assertEqual(loaded.predict(replay_rows), expected)
|
||||
independently_loaded = load_verified_model_bundle(right)
|
||||
self.assertEqual(
|
||||
loaded.prediction_bytes(replay_rows),
|
||||
independently_loaded.prediction_bytes(replay_rows),
|
||||
)
|
||||
self.assertEqual(loaded.canonical_bytes, left.read_bytes())
|
||||
|
||||
def test_content_tampering_fails_closed(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
path = write_model_bundle(
|
||||
self.bundle,
|
||||
Path(directory) / "bundle.json",
|
||||
)
|
||||
tampered = json.loads(path.read_text(encoding="utf-8"))
|
||||
tampered["model"]["coefficients"][0] += 0.01
|
||||
path.write_text(
|
||||
canonical_json(tampered) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
with self.assertRaisesRegex(
|
||||
ModelBundleIntegrityError,
|
||||
"SHA-256 mismatch",
|
||||
):
|
||||
load_verified_model_bundle(path)
|
||||
|
||||
def test_resealed_semantic_tampering_fails_closed(self):
|
||||
tampered = copy.deepcopy(self.bundle)
|
||||
tampered["preprocessor"]["output_feature_order"] = list(
|
||||
reversed(tampered["preprocessor"]["output_feature_order"])
|
||||
)
|
||||
tampered["bundle_sha256"] = model_bundle_sha256(tampered)
|
||||
with self.assertRaisesRegex(
|
||||
ModelBundleIntegrityError,
|
||||
"output_feature_order is inconsistent",
|
||||
):
|
||||
verify_model_bundle(tampered)
|
||||
|
||||
def test_expected_identity_rejects_a_valid_but_wrong_bundle(self):
|
||||
with self.assertRaisesRegex(
|
||||
ModelBundleIntegrityError,
|
||||
"data_sha256 does not match expected identity",
|
||||
):
|
||||
verify_model_bundle(
|
||||
self.bundle,
|
||||
expected_data_sha256="f" * 64,
|
||||
)
|
||||
|
||||
def test_placeholder_lineage_hash_is_rejected(self):
|
||||
with self.assertRaisesRegex(
|
||||
ModelBundleIntegrityError,
|
||||
"all-zero placeholder",
|
||||
):
|
||||
create_model_bundle(
|
||||
self.processor,
|
||||
self.model,
|
||||
trained_through="2024-12-31",
|
||||
label_name="label",
|
||||
horizon_trading_days=5,
|
||||
data_sha256="0" * 64,
|
||||
source_sha256=HASHES["source_sha256"],
|
||||
config_sha256=HASHES["config_sha256"],
|
||||
)
|
||||
|
||||
def test_unfitted_components_and_non_finite_state_are_rejected(self):
|
||||
with self.assertRaisesRegex(ValueError, "preprocessor is not fitted"):
|
||||
create_model_bundle(
|
||||
TrainOnlyFeaturePreprocessor(),
|
||||
self.model,
|
||||
trained_through="2024-12-31",
|
||||
label_name="label",
|
||||
horizon_trading_days=5,
|
||||
**HASHES,
|
||||
)
|
||||
tampered = copy.deepcopy(self.bundle)
|
||||
tampered["model"]["intercept"] = float("nan")
|
||||
with self.assertRaisesRegex(
|
||||
ModelBundleIntegrityError,
|
||||
"canonical JSON",
|
||||
):
|
||||
verify_model_bundle(tampered)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user