Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
91 changes: 89 additions & 2 deletions src/quant_platform_kit/strategy_lifecycle/performance_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

import hashlib
import json
import math
import tempfile
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
Expand All @@ -26,6 +27,7 @@
from quant_platform_kit.cloud import get_object_store
from quant_platform_kit.strategy_lifecycle.contracts import (
BacktestResult,
BacktestValidationIdentity,
DriftResult,
OptimizationProposal,
StrategyHealthScore,
Expand Down Expand Up @@ -673,6 +675,91 @@ def _drift_from_dict(data: Mapping[str, Any]) -> DriftResult | None:
return None


def _finite_nonnegative_number(value: object) -> float:
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError("cost_inputs")
number = float(value)
if not math.isfinite(number) or number < 0:
raise ValueError("cost_inputs")
return number


def _cost_inputs_from_payload(data: Mapping[str, Any]) -> dict[str, float]:
if "cost_inputs" not in data:
return {}
raw = data["cost_inputs"]
if not isinstance(raw, Mapping):
raise ValueError("cost_inputs")
parsed: dict[str, float] = {}
for key, value in raw.items():
if not isinstance(key, str) or not key:
raise ValueError("cost_inputs")
parsed[key] = _finite_nonnegative_number(value)
return parsed


def _stored_text(value: object) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError("validation_identity")
return value


def _optional_stored_date(value: object) -> date | None:
if value is None:
return None
if not isinstance(value, str) or not value:
raise ValueError("validation_identity")
return date.fromisoformat(value)


def _required_stored_date(value: object) -> date:
if not isinstance(value, str) or not value:
raise ValueError("validation_identity")
return date.fromisoformat(value)


def _positive_int(value: object) -> int:
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
raise ValueError("validation_identity")
return value


def _validation_identity_from_payload(data: Mapping[str, Any]) -> BacktestValidationIdentity | None:
if "validation_identity" not in data or data["validation_identity"] is None:
return None
raw = data["validation_identity"]
if not isinstance(raw, Mapping):
raise ValueError("validation_identity")
required = (
"protocol",
"fold_id",
"fold_role",
"train_start",
"train_end",
"test_start",
"test_end",
"locked_oos_start",
"locked_oos_end",
"purge_days",
"embargo_days",
)
if any(name not in raw for name in required):
raise ValueError("validation_identity")
return BacktestValidationIdentity(
protocol=_stored_text(raw["protocol"]),
fold_id=_stored_text(raw["fold_id"]),
fold_role=_stored_text(raw["fold_role"]),
train_start=_optional_stored_date(raw["train_start"]),
train_end=_optional_stored_date(raw["train_end"]),
test_start=_required_stored_date(raw["test_start"]),
test_end=_required_stored_date(raw["test_end"]),
locked_oos_start=_required_stored_date(raw["locked_oos_start"]),
locked_oos_end=_required_stored_date(raw["locked_oos_end"]),
purge_days=_positive_int(raw["purge_days"]),
embargo_days=_positive_int(raw["embargo_days"]),
)


def _backtest_from_dict(data: Mapping[str, Any]) -> BacktestResult | None:
try:
return BacktestResult(
Expand Down Expand Up @@ -710,8 +797,8 @@ def _backtest_from_dict(data: Mapping[str, Any]) -> BacktestResult | None:
computed_at=str(data.get("computed_at", "")),
source_revision=data.get("source_revision") if isinstance(data.get("source_revision"), str) else "",
cost_model=data.get("cost_model") if isinstance(data.get("cost_model"), str) else "",
validation_identity=None,
cost_inputs={},
validation_identity=_validation_identity_from_payload(data),
cost_inputs=_cost_inputs_from_payload(data),
periods_per_year=(
float(data["periods_per_year"]) if data.get("periods_per_year") is not None else None
),
Expand Down
98 changes: 97 additions & 1 deletion tests/test_lifecycle_performance_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,10 +6,14 @@
import json
import tempfile
import unittest
from datetime import date
from pathlib import Path
from unittest.mock import patch

from quant_platform_kit.strategy_lifecycle.contracts import BacktestResult
from quant_platform_kit.strategy_lifecycle.contracts import (
BacktestResult,
BacktestValidationIdentity,
)
from quant_platform_kit.strategy_lifecycle.performance_store import (
DEFAULT_LOCAL_ROOT,
PerformanceStore,
Expand Down Expand Up @@ -490,5 +494,97 @@ def test_same_run_merges_newer_legacy_with_new_path_and_pins_version(self) -> No
self.assertEqual(pinned_new.sharpe_ratio, 2.0)


def _validation_identity() -> BacktestValidationIdentity:
return BacktestValidationIdentity(
protocol="purged_walk_forward.v1",
fold_id="candidate_wf0",
fold_role="test",
train_start=date(2020, 1, 2),
train_end=date(2021, 1, 4),
test_start=date(2021, 2, 1),
test_end=date(2021, 6, 1),
locked_oos_start=date(2022, 1, 3),
locked_oos_end=date(2023, 1, 4),
purge_days=5,
embargo_days=3,
)


def _plant_backtest(root: Path, run_id: str, payload: dict[str, object]) -> None:
path = _run_file(root, run_id)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload), encoding="utf-8")


class BacktestReadbackMetadataTest(unittest.TestCase):
def test_exact_readback_keeps_two_distinct_cost_inputs(self) -> None:
low = {"commission_bps": 1.0, "slippage_bps": 2.0, "market_impact_bps": 0.0}
high = {"commission_bps": 5.5, "slippage_bps": 3.0, "market_impact_bps": 1.25}
with tempfile.TemporaryDirectory() as tmp:
store = PerformanceStore(local_root=Path(tmp))
store.save_backtest_result(_backtest(run_id="cost-low", cost_inputs=low, sharpe_ratio=0.4))
store.save_backtest_result(_backtest(run_id="cost-high", cost_inputs=high, sharpe_ratio=0.8))
loaded_low = store.load_backtest_by_run_id("us_equity", "global_etf_rotation", "cost-low")
loaded_high = store.load_backtest_by_run_id("us_equity", "global_etf_rotation", "cost-high")

self.assertEqual(dict(loaded_low.cost_inputs), low)
self.assertEqual(dict(loaded_high.cost_inputs), high)
self.assertIsNone(loaded_low.validation_identity)
self.assertIsNone(loaded_high.validation_identity)

def test_validation_identity_round_trip(self) -> None:
identity = _validation_identity()
costs = {"commission_bps": 2.0, "slippage_bps": 1.0, "market_impact_bps": 0.5}
with tempfile.TemporaryDirectory() as tmp:
store = PerformanceStore(local_root=Path(tmp))
store.save_backtest_result(_backtest(
run_id="identity-run",
validation_identity=identity,
cost_inputs=costs,
))
loaded = store.load_backtest_by_run_id("us_equity", "global_etf_rotation", "identity-run")

self.assertEqual(loaded.validation_identity, identity)
self.assertEqual(dict(loaded.cost_inputs), costs)

def test_legacy_file_without_cost_or_identity_keeps_defaults(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
store = PerformanceStore(local_root=root)
payload = _backtest(run_id="legacy-meta", sharpe_ratio=0.2).to_dict()
payload.pop("cost_inputs")
payload.pop("validation_identity")
_plant_backtest(root, "legacy-meta", payload)
loaded = store.load_backtest_by_run_id("us_equity", "global_etf_rotation", "legacy-meta")

self.assertIsNotNone(loaded)
self.assertEqual(dict(loaded.cost_inputs), {})
self.assertIsNone(loaded.validation_identity)
self.assertEqual(loaded.sharpe_ratio, 0.2)

def test_malformed_cost_inputs_or_identity_fail_closed(self) -> None:
identity = _validation_identity().to_dict()
malformed: dict[str, dict[str, object]] = {
"null-cost": {"cost_inputs": None},
"bool-cost": {"cost_inputs": {"commission_bps": True}},
"text-cost": {"cost_inputs": {"commission_bps": "1"}},
"nonfinite-cost": {"cost_inputs": {"commission_bps": float("nan")}},
"negative-cost": {"cost_inputs": {"commission_bps": -1}},
"bad-identity-type": {"validation_identity": "purged_walk_forward.v1"},
"partial-identity": {"validation_identity": {"protocol": "purged_walk_forward.v1"}},
"bool-purge": {"validation_identity": {**identity, "purge_days": True}},
}
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
store = PerformanceStore(local_root=root)
for run_id, changes in malformed.items():
payload = _backtest(run_id=run_id, cost_inputs={"commission_bps": 1.0}).to_dict()
payload["validation_identity"] = identity
payload.update(changes)
_plant_backtest(root, run_id, payload)
loaded = store.load_backtest_by_run_id("us_equity", "global_etf_rotation", run_id)
self.assertIsNone(loaded, run_id)


if __name__ == "__main__":
unittest.main()
Loading