Skip to content

Commit a306ab9

Browse files
Pigbibicodex
andcommitted
fix: harden accepted baseline drift handling
Co-Authored-By: Codex <noreply@openai.com>
1 parent c5556bf commit a306ab9

9 files changed

Lines changed: 217 additions & 29 deletions

File tree

src/quant_platform_kit/strategy_lifecycle/cli.py

Lines changed: 66 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
import importlib
77
import sys
88
from collections.abc import Callable, Sequence
9-
from dataclasses import replace
9+
from pathlib import Path
1010
from typing import Any
1111

1212

@@ -34,24 +34,71 @@ def _run_monitor(args: argparse.Namespace) -> int:
3434
return 0
3535

3636

37+
def _parse_baseline_bucket(value: str) -> tuple[str, str]:
38+
location = value.strip()
39+
if location.startswith("gs://"):
40+
location = location[5:]
41+
bucket, _, prefix = location.partition("/")
42+
if not bucket:
43+
raise ValueError("baseline bucket must include a bucket name")
44+
return bucket, prefix.strip("/")
45+
46+
47+
def _baseline_store_from_args(args: argparse.Namespace):
48+
from quant_platform_kit.strategy_lifecycle.performance_store import PerformanceStore
49+
50+
local_root = getattr(args, "baseline_local_root", None)
51+
bucket_value = getattr(args, "baseline_bucket", None)
52+
if not local_root and not bucket_value:
53+
return None
54+
if not bucket_value:
55+
return PerformanceStore(local_root=Path(local_root))
56+
57+
bucket, prefix = _parse_baseline_bucket(bucket_value)
58+
environment_store = PerformanceStore.from_env()
59+
return PerformanceStore(
60+
cloud_bucket=bucket,
61+
cloud_prefix=prefix,
62+
local_root=Path(local_root) if local_root else None,
63+
project_id=environment_store.project_id,
64+
client_factory=environment_store.client_factory,
65+
)
66+
67+
68+
def _baseline_lineage_policy_from_args(args: argparse.Namespace) -> str:
69+
allow_legacy = getattr(args, "allow_legacy_baseline_history", False)
70+
strict = getattr(args, "strict_baseline_lineage", False)
71+
if allow_legacy and strict:
72+
raise ValueError("baseline lineage flags are mutually exclusive")
73+
if allow_legacy:
74+
return "migration"
75+
return "strict" if strict else "auto"
76+
77+
78+
def _add_baseline_options(parser: argparse.ArgumentParser) -> None:
79+
parser.add_argument("--baseline-local-root", default=None)
80+
parser.add_argument(
81+
"--baseline-bucket",
82+
default=None,
83+
help="Accepted-baseline bucket or gs://bucket/prefix URI.",
84+
)
85+
lineage = parser.add_mutually_exclusive_group()
86+
lineage.add_argument("--strict-baseline-lineage", action="store_true")
87+
lineage.add_argument(
88+
"--allow-legacy-baseline-history",
89+
action="store_true",
90+
help="One-time migration: reuse untagged prior drift status with an external accepted baseline.",
91+
)
92+
93+
3794
def _run_drift(args: argparse.Namespace) -> int:
3895
_print(f"[drift] Running drift detection for domain={args.domain}")
3996
run_drift_detection = _load_callable(
4097
"quant_platform_kit.strategy_lifecycle.drift_detector",
4198
"run_drift_detection",
4299
)
43-
baseline_store = None
44-
if getattr(args, "baseline_local_root", None):
45-
from pathlib import Path
46-
from quant_platform_kit.strategy_lifecycle.performance_store import PerformanceStore
47-
48-
baseline_store = replace(
49-
PerformanceStore.from_env(),
50-
local_root=Path(args.baseline_local_root),
51-
)
52-
baseline_lineage_policy = "migration" if getattr(args, "allow_legacy_baseline_history", False) else (
53-
"strict" if getattr(args, "strict_baseline_lineage", False) else "auto"
54-
)
100+
baseline_store = _baseline_store_from_args(args)
101+
baseline_lineage_policy = _baseline_lineage_policy_from_args(args)
55102
results = run_drift_detection(
56103
domain=args.domain,
57104
strategy_profile=args.strategy,
@@ -153,6 +200,10 @@ def _run_lifecycle(args: argparse.Namespace) -> int:
153200
strategy=None,
154201
no_alerts=args.no_alerts,
155202
dry_run_alerts=args.dry_run_alerts,
203+
baseline_local_root=getattr(args, "baseline_local_root", None),
204+
baseline_bucket=getattr(args, "baseline_bucket", None),
205+
strict_baseline_lineage=getattr(args, "strict_baseline_lineage", False),
206+
allow_legacy_baseline_history=getattr(args, "allow_legacy_baseline_history", False),
156207
)
157208
)
158209
if drift_status != 0:
@@ -251,13 +302,7 @@ def build_parser() -> argparse.ArgumentParser:
251302
drift.add_argument("--strategy", default=None)
252303
drift.add_argument("--no-alerts", action="store_true")
253304
drift.add_argument("--dry-run-alerts", action="store_true")
254-
drift.add_argument("--baseline-local-root", default=None)
255-
drift.add_argument("--strict-baseline-lineage", action="store_true")
256-
drift.add_argument(
257-
"--allow-legacy-baseline-history",
258-
action="store_true",
259-
help="One-time migration: reuse untagged prior drift status with an external accepted baseline.",
260-
)
305+
_add_baseline_options(drift)
261306
drift.set_defaults(func=_run_drift)
262307

263308
optimize = subparsers.add_parser("optimize", help="Run parameter optimization for one strategy.")
@@ -309,6 +354,7 @@ def build_parser() -> argparse.ArgumentParser:
309354
lifecycle.add_argument("--skip-optimization", action="store_true")
310355
lifecycle.add_argument("--no-alerts", action="store_true")
311356
lifecycle.add_argument("--dry-run-alerts", action="store_true")
357+
_add_baseline_options(lifecycle)
312358
lifecycle.set_defaults(func=_run_lifecycle)
313359

314360
return parser

src/quant_platform_kit/strategy_lifecycle/codex_integration.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -403,7 +403,7 @@ def _run_drift_phase(domain: str, store: PerformanceStore) -> tuple[list, list]:
403403
"""Phase 2: run drift detection, return (all_drifts, alerting_drifts)."""
404404
from quant_platform_kit.strategy_lifecycle.drift_detector import run_drift_detection
405405
drifts = run_drift_detection(domain, store=store)
406-
alerts = [d for d in drifts if d.status != DriftStatus.HEALTHY]
406+
alerts = [d for d in drifts if d.status != DriftStatus.HEALTHY and not d.alert_suppressed]
407407
return drifts, alerts
408408

409409

src/quant_platform_kit/strategy_lifecycle/contracts.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -185,6 +185,8 @@ class DriftResult:
185185
alert_suppressed: bool = False
186186
baseline_param_set_id: str | None = None
187187
baseline_available: bool = True
188+
baseline_param_version: int | None = None
189+
baseline_artifact_id: str | None = None
188190

189191
def to_dict(self) -> dict[str, object]:
190192
return {
@@ -197,6 +199,8 @@ def to_dict(self) -> dict[str, object]:
197199
"previous_status": self.previous_status.value if self.previous_status else None,
198200
"baseline_param_set_id": self.baseline_param_set_id,
199201
"baseline_available": self.baseline_available,
202+
"baseline_param_version": self.baseline_param_version,
203+
"baseline_artifact_id": self.baseline_artifact_id,
200204
"escalated": self.escalated,
201205
"cooldown_active": self.cooldown_active,
202206
"alert_suppressed": self.alert_suppressed,

src/quant_platform_kit/strategy_lifecycle/drift_detector.py

Lines changed: 42 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,12 @@
3434
]
3535

3636

37+
def _baseline_artifact_id(backtest: BacktestResult | None) -> str | None:
38+
if backtest is None:
39+
return None
40+
return backtest.run_id or backtest.computed_at or None
41+
42+
3743
def _compute_dimension(
3844
key: str, metric: str,
3945
actual_val: float, expected_val: float, threshold: float,
@@ -127,6 +133,8 @@ def detect_drift(
127133
status=status, dimensions=dimensions,
128134
previous_status=previous_status,
129135
baseline_param_set_id=backtest.param_set_id if backtest else None,
136+
baseline_param_version=backtest.param_version if backtest else None,
137+
baseline_artifact_id=_baseline_artifact_id(backtest),
130138
escalated=escalated,
131139
)
132140

@@ -149,7 +157,7 @@ def run_drift_detection(
149157
it writes the accepted baseline ID into the next result.
150158
"""
151159
store = store or PerformanceStore.from_env()
152-
explicit_baseline_store = baseline_store is not None
160+
explicit_baseline_store = baseline_store is not None and baseline_store is not store
153161
if baseline_lineage_policy not in {"auto", "compatible", "migration", "strict"}:
154162
raise ValueError("baseline_lineage_policy must be auto, compatible, migration, or strict")
155163
if explicit_baseline_store and baseline_lineage_policy == "compatible":
@@ -185,14 +193,34 @@ def run_drift_detection(
185193
previous = read_previous.load_latest_drift(domain, profile)
186194
previous_before_lineage_check = previous
187195
current_baseline_id = backtest.param_set_id if backtest else None
196+
current_baseline_version = backtest.param_version if backtest else None
197+
current_baseline_artifact_id = _baseline_artifact_id(backtest)
188198
if previous:
189199
previous_baseline_id = previous.baseline_param_set_id
200+
previous_baseline_version = previous.baseline_param_version
201+
previous_baseline_artifact_id = previous.baseline_artifact_id
190202
if baseline_lineage_policy == "strict":
191-
if not (previous_baseline_id and current_baseline_id and previous_baseline_id == current_baseline_id):
203+
if not (
204+
previous_baseline_id
205+
and current_baseline_id
206+
and previous_baseline_id == current_baseline_id
207+
and previous_baseline_version is not None
208+
and previous_baseline_version == current_baseline_version
209+
and (
210+
not current_baseline_artifact_id
211+
or previous_baseline_artifact_id == current_baseline_artifact_id
212+
)
213+
):
192214
previous = None
193215
elif baseline_lineage_policy == "migration":
194216
if current_baseline_id is None or (
195217
previous_baseline_id is not None and previous_baseline_id != current_baseline_id
218+
) or (
219+
previous_baseline_version is not None
220+
and previous_baseline_version != current_baseline_version
221+
) or (
222+
previous_baseline_artifact_id is not None
223+
and previous_baseline_artifact_id != current_baseline_artifact_id
196224
):
197225
previous = None
198226
if backtest is None:
@@ -213,6 +241,18 @@ def run_drift_detection(
213241
previous_status=continuity_result.status,
214242
alert_suppressed=True,
215243
baseline_param_set_id=continuity_result.baseline_param_set_id,
244+
baseline_param_version=continuity_result.baseline_param_version,
245+
baseline_artifact_id=continuity_result.baseline_artifact_id,
246+
baseline_available=False,
247+
)
248+
elif explicit_baseline_store:
249+
result = DriftResult(
250+
strategy_profile=snapshot.strategy_profile,
251+
domain=snapshot.domain,
252+
as_of=snapshot.as_of,
253+
drift_score=0.0,
254+
status=DriftStatus.REVIEW,
255+
alert_suppressed=True,
216256
baseline_available=False,
217257
)
218258
else:

src/quant_platform_kit/strategy_lifecycle/performance_store.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -538,6 +538,16 @@ def _drift_from_dict(data: Mapping[str, Any]) -> DriftResult | None:
538538
previous_status=DriftStatus(str(data["previous_status"])) if data.get("previous_status") else None,
539539
baseline_param_set_id=str(data["baseline_param_set_id"]) if data.get("baseline_param_set_id") else None,
540540
baseline_available=bool(data.get("baseline_available", True)),
541+
baseline_param_version=(
542+
int(data["baseline_param_version"])
543+
if data.get("baseline_param_version") is not None
544+
else None
545+
),
546+
baseline_artifact_id=(
547+
str(data["baseline_artifact_id"])
548+
if data.get("baseline_artifact_id")
549+
else None
550+
),
541551
)
542552
except Exception:
543553
return None

tests/test_lifecycle_cli.py

Lines changed: 45 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import unittest
6+
from pathlib import Path
67
from types import SimpleNamespace
78
from unittest.mock import patch
89

@@ -120,10 +121,7 @@ def fake_run_drift_detection(**kwargs):
120121
return lambda _events, **_kwargs: {}
121122
raise AssertionError(function_name)
122123

123-
environment_store = PerformanceStore(
124-
cloud_bucket="accepted-bucket",
125-
cloud_prefix="lifecycle",
126-
)
124+
environment_store = PerformanceStore(cloud_bucket="candidate-bucket", cloud_prefix="candidate")
127125
with (
128126
patch.object(cli, "_load_callable", fake_load_callable),
129127
patch.object(PerformanceStore, "from_env", return_value=environment_store),
@@ -137,8 +135,33 @@ def fake_run_drift_detection(**kwargs):
137135

138136
self.assertEqual(result, 0)
139137
self.assertEqual(observed["baseline_lineage_policy"], "migration")
138+
self.assertEqual(observed["baseline_store"].local_root, Path("accepted-baselines"))
139+
self.assertEqual(observed["baseline_store"].cloud_bucket, "")
140+
self.assertEqual(observed["baseline_store"].cloud_prefix, "")
141+
142+
def test_drift_command_uses_explicit_baseline_bucket(self) -> None:
143+
observed = {}
144+
145+
def fake_load_callable(_module_name: str, function_name: str):
146+
if function_name == "run_drift_detection":
147+
def fake_run_drift_detection(**kwargs):
148+
observed.update(kwargs)
149+
return []
150+
151+
return fake_run_drift_detection
152+
if function_name == "build_drift_alert":
153+
return lambda _result: None
154+
if function_name == "publish_drift_alerts":
155+
return lambda _events, **_kwargs: {}
156+
raise AssertionError(function_name)
157+
158+
with patch.object(cli, "_load_callable", fake_load_callable):
159+
result = cli.main(["drift", "--baseline-bucket", "gs://accepted-bucket/lifecycle"])
160+
161+
self.assertEqual(result, 0)
140162
self.assertEqual(observed["baseline_store"].cloud_bucket, "accepted-bucket")
141163
self.assertEqual(observed["baseline_store"].cloud_prefix, "lifecycle")
164+
self.assertIsNone(observed["baseline_store"].local_root)
142165

143166
def test_update_returns_non_zero_for_error_stage(self) -> None:
144167
def fake_load_callable(_module_name: str, _function_name: str):
@@ -200,6 +223,7 @@ def fake_load_callable(_module_name: str, _function_name: str):
200223

201224
def test_lifecycle_command_runs_real_steps(self) -> None:
202225
calls = []
226+
drift_kwargs = {}
203227

204228
def fake_load_callable(_module_name: str, function_name: str):
205229
if function_name == "run_monitor":
@@ -209,8 +233,9 @@ def fake_monitor(**_kwargs):
209233

210234
return fake_monitor
211235
if function_name == "run_drift_detection":
212-
def fake_drift(**_kwargs):
236+
def fake_drift(**kwargs):
213237
calls.append("drift")
238+
drift_kwargs.update(kwargs)
214239
return []
215240

216241
return fake_drift
@@ -227,10 +252,24 @@ def fake_dashboard(**_kwargs):
227252
raise AssertionError(function_name)
228253

229254
with patch.object(cli, "_load_callable", fake_load_callable):
230-
result = cli.main(["lifecycle", "--domain", "cn_equity", "--skip-optimization"])
255+
result = cli.main([
256+
"lifecycle",
257+
"--domain",
258+
"cn_equity",
259+
"--skip-optimization",
260+
"--baseline-local-root",
261+
"accepted-baselines",
262+
"--strict-baseline-lineage",
263+
])
231264

232265
self.assertEqual(result, 0)
233266
self.assertEqual(calls, ["monitor", "drift", "dashboard"])
267+
self.assertEqual(drift_kwargs["baseline_store"].local_root, Path("accepted-baselines"))
268+
self.assertEqual(drift_kwargs["baseline_lineage_policy"], "strict")
269+
270+
def test_drift_rejects_conflicting_baseline_lineage_flags(self) -> None:
271+
with self.assertRaises(SystemExit):
272+
cli.main(["drift", "--strict-baseline-lineage", "--allow-legacy-baseline-history"])
234273

235274
def test_error_returns_non_zero(self) -> None:
236275
def fake_load_callable(_module_name: str, _function_name: str):
Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
from datetime import date
2+
from unittest.mock import Mock, patch
3+
4+
from quant_platform_kit.strategy_lifecycle.codex_integration import _run_drift_phase
5+
from quant_platform_kit.strategy_lifecycle.contracts import DriftResult, DriftStatus
6+
7+
8+
def test_drift_phase_excludes_suppressed_results_from_automation() -> None:
9+
suppressed = DriftResult(
10+
strategy_profile="missing-baseline",
11+
domain="us_equity",
12+
as_of=date(2026, 7, 11),
13+
drift_score=0.0,
14+
status=DriftStatus.REVIEW,
15+
alert_suppressed=True,
16+
baseline_available=False,
17+
)
18+
critical = DriftResult(
19+
strategy_profile="active-baseline",
20+
domain="us_equity",
21+
as_of=date(2026, 7, 11),
22+
drift_score=0.8,
23+
status=DriftStatus.CRITICAL,
24+
)
25+
26+
with patch(
27+
"quant_platform_kit.strategy_lifecycle.drift_detector.run_drift_detection",
28+
return_value=[suppressed, critical],
29+
):
30+
drifts, alerts = _run_drift_phase("us_equity", Mock())
31+
32+
assert drifts == [suppressed, critical]
33+
assert alerts == [critical]

0 commit comments

Comments
 (0)