strategy.py 8.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248
  1. from __future__ import annotations
  2. import copy
  3. import uuid
  4. from typing import Any
  5. from sqlalchemy import select
  6. from supply_infra.db.models.business_harness import (
  7. DemandEvaluationSnapshot,
  8. StrategyChangeProposal,
  9. StrategyVersion,
  10. )
  11. from supply_infra.db.session import get_session
  12. from supply_infra.pipeline.dates import china_now
  13. def _deep_merge(base: dict[str, Any], patch: dict[str, Any]) -> dict[str, Any]:
  14. result = copy.deepcopy(base)
  15. for key, value in patch.items():
  16. if isinstance(value, dict) and isinstance(result.get(key), dict):
  17. result[key] = _deep_merge(result[key], value)
  18. else:
  19. result[key] = value
  20. return result
  21. def _flatten(payload: dict[str, Any], prefix: str = "") -> dict[str, Any]:
  22. result: dict[str, Any] = {}
  23. for key, value in payload.items():
  24. path = f"{prefix}.{key}" if prefix else key
  25. if isinstance(value, dict):
  26. result.update(_flatten(value, path))
  27. else:
  28. result[path] = value
  29. return result
  30. def _validate_definition(definition: dict[str, Any]) -> None:
  31. required = {
  32. "validity_weights",
  33. "local_priority_weights",
  34. "action_thresholds",
  35. }
  36. missing = required - definition.keys()
  37. if missing:
  38. raise ValueError(f"strategy definition missing: {sorted(missing)}")
  39. for group in ("validity_weights", "local_priority_weights"):
  40. weights = definition[group]
  41. if not isinstance(weights, dict) or not weights:
  42. raise ValueError(f"{group} must be a non-empty object")
  43. for key, value in weights.items():
  44. if not isinstance(value, (int, float)) or value < 0 or value > 1:
  45. raise ValueError(f"invalid weight {group}.{key}: {value!r}")
  46. if sum(float(value) for value in weights.values()) <= 0:
  47. raise ValueError(f"{group} total weight must be positive")
  48. thresholds = _flatten(definition["action_thresholds"])
  49. for key, value in thresholds.items():
  50. if not isinstance(value, (int, float)) or not 0 <= float(value) <= 1:
  51. raise ValueError(f"invalid threshold action_thresholds.{key}: {value!r}")
  52. def _tier(definition: dict[str, Any], row: DemandEvaluationSnapshot) -> str:
  53. thresholds = definition["action_thresholds"]
  54. validity = float(row.validity_score)
  55. priority = float(row.local_supply_priority)
  56. confidence = float(row.data_confidence)
  57. assure = thresholds["assure_supply"]
  58. if (
  59. row.posterior_state == "available"
  60. and validity >= assure["validity"]
  61. and priority >= assure["priority"]
  62. ):
  63. return "assure_supply"
  64. priority_rule = thresholds["priority"]
  65. if validity >= priority_rule["validity"] and priority >= priority_rule["priority"]:
  66. return "priority"
  67. validate = thresholds["targeted_validation"]
  68. if (
  69. validity >= validate["validity"]
  70. and confidence < validate["confidence_below"]
  71. ):
  72. return "targeted_validation"
  73. if validity >= thresholds["exploration"]["validity"]:
  74. return "exploration"
  75. return "suppress"
  76. def create_strategy_proposal(
  77. *,
  78. requested_change: dict[str, Any],
  79. reason: str,
  80. created_by: str,
  81. ) -> dict[str, Any]:
  82. if not requested_change or not reason.strip() or not created_by.strip():
  83. raise ValueError("requested_change, reason and created_by are required")
  84. with get_session() as session:
  85. base = session.scalar(
  86. select(StrategyVersion)
  87. .where(StrategyVersion.status == "active")
  88. .order_by(StrategyVersion.activated_at.desc())
  89. .limit(1)
  90. )
  91. if base is None:
  92. raise RuntimeError("no active strategy version")
  93. proposed = _deep_merge(dict(base.definition_json), requested_change)
  94. _validate_definition(proposed)
  95. before = _flatten(dict(base.definition_json))
  96. after = _flatten(proposed)
  97. diff = {
  98. key: {"before": before.get(key), "after": after.get(key)}
  99. for key in sorted(set(before) | set(after))
  100. if before.get(key) != after.get(key)
  101. }
  102. proposal_id = str(uuid.uuid4())
  103. session.add(
  104. StrategyChangeProposal(
  105. proposal_id=proposal_id,
  106. base_strategy_version_id=base.strategy_version_id,
  107. status="draft",
  108. requested_change_json=requested_change,
  109. proposed_definition_json=proposed,
  110. diff_json=diff,
  111. preview_json=None,
  112. reason=reason,
  113. created_by=created_by[:128],
  114. approved_by=None,
  115. approved_at=None,
  116. )
  117. )
  118. session.flush()
  119. return {
  120. "proposal_id": proposal_id,
  121. "status": "draft",
  122. "base_strategy_version": base.version_key,
  123. "diff": diff,
  124. }
  125. def preview_strategy_proposal(proposal_id: str) -> dict[str, Any]:
  126. with get_session() as session:
  127. proposal = session.get(
  128. StrategyChangeProposal,
  129. proposal_id,
  130. with_for_update=True,
  131. )
  132. if proposal is None:
  133. raise LookupError(f"strategy proposal not found: {proposal_id}")
  134. if proposal.status not in {"draft", "previewed"}:
  135. raise RuntimeError(f"proposal cannot be previewed: {proposal.status}")
  136. evaluations = list(
  137. session.scalars(select(DemandEvaluationSnapshot)).all()
  138. )
  139. changes: dict[str, int] = {}
  140. changed_items: list[dict[str, Any]] = []
  141. definition = dict(proposal.proposed_definition_json)
  142. for row in evaluations:
  143. new_tier = _tier(definition, row)
  144. if new_tier == row.action_tier:
  145. continue
  146. key = f"{row.action_tier}->{new_tier}"
  147. changes[key] = changes.get(key, 0) + 1
  148. changed_items.append(
  149. {
  150. "evaluation_id": row.evaluation_id,
  151. "biz_dt": row.biz_dt,
  152. "before": row.action_tier,
  153. "after": new_tier,
  154. }
  155. )
  156. preview = {
  157. "evaluated_count": len(evaluations),
  158. "changed_count": len(changed_items),
  159. "tier_changes": changes,
  160. "changed_items": changed_items[:500],
  161. "truncated": len(changed_items) > 500,
  162. }
  163. proposal.preview_json = preview
  164. proposal.status = "previewed"
  165. session.flush()
  166. return {
  167. "proposal_id": proposal_id,
  168. "status": proposal.status,
  169. "preview": preview,
  170. }
  171. def approve_strategy_proposal(
  172. proposal_id: str,
  173. *,
  174. approved_by: str,
  175. ) -> dict[str, Any]:
  176. if not approved_by.strip():
  177. raise ValueError("approved_by is required")
  178. with get_session() as session:
  179. proposal = session.get(
  180. StrategyChangeProposal,
  181. proposal_id,
  182. with_for_update=True,
  183. )
  184. if proposal is None:
  185. raise LookupError(f"strategy proposal not found: {proposal_id}")
  186. if proposal.status == "approved":
  187. version = session.scalar(
  188. select(StrategyVersion).where(
  189. StrategyVersion.change_reason
  190. == f"approved proposal {proposal_id}: {proposal.reason}"
  191. )
  192. )
  193. return {
  194. "proposal_id": proposal_id,
  195. "status": "approved",
  196. "strategy_version": version.version_key if version else None,
  197. "idempotent_replay": True,
  198. }
  199. if proposal.status != "previewed" or proposal.preview_json is None:
  200. raise RuntimeError("proposal must be previewed before approval")
  201. active_rows = list(
  202. session.scalars(
  203. select(StrategyVersion)
  204. .where(StrategyVersion.status == "active")
  205. .with_for_update()
  206. ).all()
  207. )
  208. for row in active_rows:
  209. row.status = "retired"
  210. version_key = f"harness-{proposal_id}"
  211. version = StrategyVersion(
  212. strategy_version_id=str(uuid.uuid4()),
  213. version_key=version_key,
  214. status="active",
  215. definition_json=proposal.proposed_definition_json,
  216. change_reason=f"approved proposal {proposal_id}: {proposal.reason}",
  217. created_by=approved_by[:128],
  218. activated_at=china_now(),
  219. )
  220. session.add(version)
  221. proposal.status = "approved"
  222. proposal.approved_by = approved_by[:128]
  223. proposal.approved_at = china_now()
  224. session.flush()
  225. return {
  226. "proposal_id": proposal_id,
  227. "status": proposal.status,
  228. "strategy_version": version.version_key,
  229. "strategy_version_id": version.strategy_version_id,
  230. "idempotent_replay": False,
  231. }