| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248 |
- 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,
- }
|