from __future__ import annotations import copy import uuid from typing import Any from sqlalchemy import select from supply_infra.db.models.business_harness import ( DemandEvaluationSnapshot, StrategyChangeProposal, StrategyVersion, ) from supply_infra.db.session import get_session from supply_infra.pipeline.dates import china_now def _deep_merge(base: dict[str, Any], patch: dict[str, Any]) -> dict[str, Any]: result = copy.deepcopy(base) for key, value in patch.items(): if isinstance(value, dict) and isinstance(result.get(key), dict): result[key] = _deep_merge(result[key], value) else: result[key] = value return result def _flatten(payload: dict[str, Any], prefix: str = "") -> dict[str, Any]: result: dict[str, Any] = {} for key, value in payload.items(): path = f"{prefix}.{key}" if prefix else key if isinstance(value, dict): result.update(_flatten(value, path)) else: result[path] = value return result def _validate_definition(definition: dict[str, Any]) -> None: required = { "validity_weights", "local_priority_weights", "action_thresholds", } missing = required - definition.keys() if missing: raise ValueError(f"strategy definition missing: {sorted(missing)}") for group in ("validity_weights", "local_priority_weights"): weights = definition[group] if not isinstance(weights, dict) or not weights: raise ValueError(f"{group} must be a non-empty object") for key, value in weights.items(): if not isinstance(value, (int, float)) or value < 0 or value > 1: raise ValueError(f"invalid weight {group}.{key}: {value!r}") if sum(float(value) for value in weights.values()) <= 0: raise ValueError(f"{group} total weight must be positive") thresholds = _flatten(definition["action_thresholds"]) for key, value in thresholds.items(): if not isinstance(value, (int, float)) or not 0 <= float(value) <= 1: raise ValueError(f"invalid threshold action_thresholds.{key}: {value!r}") def _tier(definition: dict[str, Any], row: DemandEvaluationSnapshot) -> str: thresholds = definition["action_thresholds"] validity = float(row.validity_score) priority = float(row.local_supply_priority) confidence = float(row.data_confidence) assure = thresholds["assure_supply"] if ( row.posterior_state == "available" and validity >= assure["validity"] and priority >= assure["priority"] ): return "assure_supply" priority_rule = thresholds["priority"] if validity >= priority_rule["validity"] and priority >= priority_rule["priority"]: return "priority" validate = thresholds["targeted_validation"] if ( validity >= validate["validity"] and confidence < validate["confidence_below"] ): return "targeted_validation" if validity >= thresholds["exploration"]["validity"]: return "exploration" return "suppress" def create_strategy_proposal( *, requested_change: dict[str, Any], reason: str, created_by: str, ) -> dict[str, Any]: if not requested_change or not reason.strip() or not created_by.strip(): raise ValueError("requested_change, reason and created_by are required") with get_session() as session: base = session.scalar( select(StrategyVersion) .where(StrategyVersion.status == "active") .order_by(StrategyVersion.activated_at.desc()) .limit(1) ) if base is None: raise RuntimeError("no active strategy version") proposed = _deep_merge(dict(base.definition_json), requested_change) _validate_definition(proposed) before = _flatten(dict(base.definition_json)) after = _flatten(proposed) diff = { key: {"before": before.get(key), "after": after.get(key)} for key in sorted(set(before) | set(after)) if before.get(key) != after.get(key) } proposal_id = str(uuid.uuid4()) session.add( StrategyChangeProposal( proposal_id=proposal_id, base_strategy_version_id=base.strategy_version_id, status="draft", requested_change_json=requested_change, proposed_definition_json=proposed, diff_json=diff, preview_json=None, reason=reason, created_by=created_by[:128], approved_by=None, approved_at=None, ) ) session.flush() return { "proposal_id": proposal_id, "status": "draft", "base_strategy_version": base.version_key, "diff": diff, } def preview_strategy_proposal(proposal_id: str) -> dict[str, Any]: with get_session() as session: proposal = session.get( StrategyChangeProposal, proposal_id, with_for_update=True, ) if proposal is None: raise LookupError(f"strategy proposal not found: {proposal_id}") if proposal.status not in {"draft", "previewed"}: raise RuntimeError(f"proposal cannot be previewed: {proposal.status}") evaluations = list( session.scalars(select(DemandEvaluationSnapshot)).all() ) changes: dict[str, int] = {} changed_items: list[dict[str, Any]] = [] definition = dict(proposal.proposed_definition_json) for row in evaluations: new_tier = _tier(definition, row) if new_tier == row.action_tier: continue key = f"{row.action_tier}->{new_tier}" changes[key] = changes.get(key, 0) + 1 changed_items.append( { "evaluation_id": row.evaluation_id, "biz_dt": row.biz_dt, "before": row.action_tier, "after": new_tier, } ) preview = { "evaluated_count": len(evaluations), "changed_count": len(changed_items), "tier_changes": changes, "changed_items": changed_items[:500], "truncated": len(changed_items) > 500, } proposal.preview_json = preview proposal.status = "previewed" session.flush() return { "proposal_id": proposal_id, "status": proposal.status, "preview": preview, } def approve_strategy_proposal( proposal_id: str, *, approved_by: str, ) -> dict[str, Any]: if not approved_by.strip(): raise ValueError("approved_by is required") with get_session() as session: proposal = session.get( StrategyChangeProposal, proposal_id, with_for_update=True, ) if proposal is None: raise LookupError(f"strategy proposal not found: {proposal_id}") if proposal.status == "approved": version = session.scalar( select(StrategyVersion).where( StrategyVersion.change_reason == f"approved proposal {proposal_id}: {proposal.reason}" ) ) return { "proposal_id": proposal_id, "status": "approved", "strategy_version": version.version_key if version else None, "idempotent_replay": True, } if proposal.status != "previewed" or proposal.preview_json is None: raise RuntimeError("proposal must be previewed before approval") active_rows = list( session.scalars( select(StrategyVersion) .where(StrategyVersion.status == "active") .with_for_update() ).all() ) for row in active_rows: row.status = "retired" version_key = f"harness-{proposal_id}" version = StrategyVersion( strategy_version_id=str(uuid.uuid4()), version_key=version_key, status="active", definition_json=proposal.proposed_definition_json, change_reason=f"approved proposal {proposal_id}: {proposal.reason}", created_by=approved_by[:128], activated_at=china_now(), ) session.add(version) proposal.status = "approved" proposal.approved_by = approved_by[:128] proposal.approved_at = china_now() session.flush() return { "proposal_id": proposal_id, "status": proposal.status, "strategy_version": version.version_key, "strategy_version_id": version.strategy_version_id, "idempotent_replay": False, }