diff --git a/src/quant_platform_kit/risk/gate.py b/src/quant_platform_kit/risk/gate.py index 08901ed..36a9432 100644 --- a/src/quant_platform_kit/risk/gate.py +++ b/src/quant_platform_kit/risk/gate.py @@ -235,6 +235,17 @@ def _snapshot_metrics( }, observed, total_equity, set() +def _completed_session_equity(portfolio_snapshot: Any) -> float | None: + if isinstance(portfolio_snapshot, Mapping): + value = portfolio_snapshot.get("completed_session_equity") + elif isinstance(portfolio_snapshot, PortfolioSnapshot): + value = portfolio_snapshot.metadata.get("completed_session_equity") + else: + return None + completed_equity = _finite_number(value) + return completed_equity if completed_equity is not None and completed_equity > 0.0 else None + + def _exact_numeric_mapping(value: Any, expected: Mapping[str, float]) -> bool: if not isinstance(value, Mapping) or set(value) != set(expected): return False @@ -867,7 +878,9 @@ def _global_etf_risk_control_fields( errors.add("invalid_risk_control_state") elif (age := (now - as_of).total_seconds()) < 0.0 or age > max_age: errors.add("stale_risk_control_state") - if risk_control_state.get("mandate_id") != mandate.get("mandate_id"): + raw_mandate_id = risk_control_state.get("mandate_id") + mandate_id = raw_mandate_id if isinstance(raw_mandate_id, str) else None + if mandate_id != mandate.get("mandate_id"): errors.add("risk_control_mandate_mismatch") if _sha256(risk_control_state.get("candidate_identity_sha256")) != mandate.get( "candidate_identity_sha256" @@ -875,7 +888,11 @@ def _global_etf_risk_control_fields( errors.add("risk_control_candidate_mismatch") if stop_loss_distance != 0.05: errors.add("invalid_stop_loss_distance") - if risk_control_state.get("stop_fill_policy") != _GLOBAL_ETF_STOP_FILL_POLICY: + raw_stop_fill_policy = risk_control_state.get("stop_fill_policy") + stop_fill_policy = ( + raw_stop_fill_policy if isinstance(raw_stop_fill_policy, str) else None + ) + if stop_fill_policy != _GLOBAL_ETF_STOP_FILL_POLICY: errors.add("invalid_stop_fill_policy") if account_drawdown is None or not 0.0 <= account_drawdown <= 1.0: errors.add("invalid_account_drawdown") @@ -951,12 +968,12 @@ def _global_etf_risk_control_fields( payload = { "as_of": _utc_timestamp(as_of) if as_of is not None else None, - "mandate_id": risk_control_state.get("mandate_id"), + "mandate_id": mandate_id, "candidate_identity_sha256": _sha256( risk_control_state.get("candidate_identity_sha256") ), "stop_loss_distance": stop_loss_distance, - "stop_fill_policy": risk_control_state.get("stop_fill_policy"), + "stop_fill_policy": stop_fill_policy, "position_stop_states": normalized_stop_states, "consecutive_completed_losing_exits": losses, "account_drawdown_fraction": account_drawdown, @@ -1004,6 +1021,12 @@ def assess_with_evidence( max_snapshot_age_seconds=mandate.get("max_snapshot_age_seconds"), ) reason_codes.update(snapshot_errors) + completed_session_equity: float | None = None + if mandate.get("mandate_id") == _GLOBAL_ETF_RESEARCH_MANDATE: + completed_session_equity = _completed_session_equity(portfolio_snapshot) + snapshot_payload["completed_session_equity"] = completed_session_equity + if completed_session_equity is None: + reason_codes.add("invalid_completed_session_equity") decision_payload, active_positions, decision_errors = _decision_metrics( decision, total_equity=total_equity, @@ -1084,15 +1107,23 @@ def assess_with_evidence( drawdown_scalar = control_fields["drawdown_scalar"] loss_budget = mandate.get("loss_budget") modeled_stop_loss = ( - sum(weight for _symbol, weight in active_positions) * stop_distance - if stop_distance is not None + sum(weight for _symbol, weight in active_positions) + * total_equity + * stop_distance + if stop_distance is not None and total_equity is not None + else None + ) + loss_budget_amount = ( + loss_budget * completed_session_equity * drawdown_scalar + if loss_budget is not None + and completed_session_equity is not None + and drawdown_scalar is not None else None ) if ( modeled_stop_loss is not None - and drawdown_scalar is not None - and loss_budget is not None - and modeled_stop_loss > loss_budget * drawdown_scalar + 1e-9 + and loss_budget_amount is not None + and modeled_stop_loss > loss_budget_amount + 1e-9 ): reason_codes.add("risk_budget_exposure_cap") target_weights: dict[str, float] = {} diff --git a/tests/test_risk_gate.py b/tests/test_risk_gate.py index 4584d7e..8163006 100644 --- a/tests/test_risk_gate.py +++ b/tests/test_risk_gate.py @@ -1356,6 +1356,7 @@ def _snapshot(**overrides: object) -> dict[str, object]: "as_of": "2026-08-09T01:59:55Z", "observed_effective_exposure": 0.0, "total_equity": 100_000.0, + "completed_session_equity": 100_000.0, } snapshot.update(overrides) return snapshot @@ -1472,6 +1473,120 @@ def test_valid_research_decision_approves_but_never_authorizes_execution( self.assertFalse(result.assessment.execution_authorized) self.assertEqual(result.decision.positions, decision.positions) + def test_loss_budget_uses_completed_session_equity_currency_basis(self) -> None: + safe_decision = _decision( + positions=(PositionTarget(symbol="XLK", target_value=20_000.0),) + ) + oversized_decision = _decision( + positions=(PositionTarget(symbol="XLK", target_value=40_000.0),) + ) + snapshot = self._snapshot( + total_equity=200_000.0, + completed_session_equity=100_000.0, + ) + risk_state = self._risk_state("XLK") + + safe, _safe_engine = self._assess( + safe_decision, + snapshot=snapshot, + risk_state=risk_state, + ) + oversized, _oversized_engine = self._assess( + oversized_decision, + snapshot=snapshot, + risk_state=risk_state, + ) + + self.assertEqual(safe.assessment.outcome, "APPROVE") + self.assertEqual(oversized.assessment.outcome, "REJECT") + self.assertIn( + "risk_budget_exposure_cap", + oversized.assessment.reason_codes, + ) + self.assertEqual(oversized.decision.positions, ()) + + def test_completed_session_equity_is_required_finite_and_digest_bound( + self, + ) -> None: + decision = _decision( + positions=(PositionTarget(symbol="XLK", target_value=10_000.0),) + ) + risk_state = self._risk_state("XLK") + invalid_values = (None, 0.0, -1.0, float("nan"), float("inf"), "100000") + for value in invalid_values: + with self.subTest(value=value): + snapshot = self._snapshot(completed_session_equity=value) + result, _engine = self._assess( + decision, + snapshot=snapshot, + risk_state=risk_state, + ) + self.assertEqual(result.assessment.outcome, "REJECT") + self.assertIn( + "invalid_completed_session_equity", + result.assessment.reason_codes, + ) + + missing_snapshot = self._snapshot() + missing_snapshot.pop("completed_session_equity") + missing, _missing_engine = self._assess( + decision, + snapshot=missing_snapshot, + risk_state=risk_state, + ) + self.assertEqual(missing.assessment.outcome, "REJECT") + self.assertIn( + "invalid_completed_session_equity", + missing.assessment.reason_codes, + ) + + first, _first_engine = self._assess( + decision, + snapshot=self._snapshot(completed_session_equity=100_000.0), + risk_state=risk_state, + ) + second, _second_engine = self._assess( + decision, + snapshot=self._snapshot(completed_session_equity=120_000.0), + risk_state=risk_state, + ) + self.assertNotEqual( + first.assessment.portfolio_snapshot_digest_sha256, + second.assessment.portfolio_snapshot_digest_sha256, + ) + + def test_malformed_risk_control_digest_material_rejects_without_raising( + self, + ) -> None: + decision = self._two_position_decision() + invalid_cases = ( + ({"stop_fill_policy": float("nan")}, "invalid_stop_fill_policy"), + ({"stop_fill_policy": float("inf")}, "invalid_stop_fill_policy"), + ({"stop_fill_policy": object()}, "invalid_stop_fill_policy"), + ({"mandate_id": float("nan")}, "risk_control_mandate_mismatch"), + ({"mandate_id": float("inf")}, "risk_control_mandate_mismatch"), + ({"mandate_id": object()}, "risk_control_mandate_mismatch"), + ({"stop_loss_distance": float("nan")}, "invalid_stop_loss_distance"), + ( + {"account_drawdown_fraction": float("inf")}, + "invalid_account_drawdown", + ), + ({"drawdown_scalar": float("nan")}, "drawdown_scalar_mismatch"), + ( + {"consecutive_completed_losing_exits": float("inf")}, + "invalid_strategy_breaker_state", + ), + ) + for overrides, reason in invalid_cases: + with self.subTest(overrides=overrides): + result, _engine = self._assess( + decision, + risk_state=self._risk_state("XLK", "BIL", **overrides), + ) + self.assertEqual(result.assessment.outcome, "REJECT") + self.assertIn(reason, result.assessment.reason_codes) + self.assertEqual(result.decision.positions, ()) + def test_exact_mandate_shape_is_fail_closed(self) -> None: decision = self._two_position_decision() caps = {symbol: 0.50 for symbol in self._ALLOWED_ASSETS}