run.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158
  1. """运行统筹规划 Agent(save_grade_plan 工具内直接入库)。"""
  2. from __future__ import annotations
  3. import json
  4. import logging
  5. from typing import Any
  6. from agents.demand_grade_orchestrator_agent import create_demand_grade_orchestrator_agent
  7. from agents.demand_grade_orchestrator_agent.common.assignment import MAX_DAILY_BATCHES, resolve_planning_state
  8. from supply_agent.types import Role
  9. from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
  10. from supply_infra.db.session import get_session
  11. logger = logging.getLogger(__name__)
  12. def _summarize_agent_saves(result: Any) -> dict[str, Any]:
  13. """统计 Agent 会话中 save_grade_plan 成功入库的次数与批次数。"""
  14. save_count = 0
  15. persisted_groups = 0
  16. last_payload: dict[str, Any] | None = None
  17. for message in result.messages:
  18. if message.role != Role.TOOL or message.name != "save_grade_plan":
  19. continue
  20. try:
  21. payload = json.loads(message.content or "")
  22. except json.JSONDecodeError:
  23. continue
  24. if not isinstance(payload, dict):
  25. continue
  26. if payload.get("ok") is True and payload.get("persisted") is True:
  27. save_count += 1
  28. persisted_groups += int(payload.get("persisted_group_count") or 0)
  29. last_payload = payload
  30. return {
  31. "save_count": save_count,
  32. "persisted_groups": persisted_groups,
  33. "last_payload": last_payload,
  34. }
  35. def _orchestrate_via_agent(
  36. biz_dt: str,
  37. *,
  38. planning_state: dict[str, Any],
  39. unassigned_count: int,
  40. ) -> None:
  41. remaining_quota = int(planning_state["remaining_batch_quota"])
  42. agent = create_demand_grade_orchestrator_agent()
  43. try:
  44. result = agent.run(
  45. f"""请为业务日 {biz_dt} 制定全局树需求分级计划。
  46. 必须从 query_global_heat_tree 开始,经过至少一次 query_heat_node_group 下钻,
  47. 由你自行划分批次后调用 save_grade_plan 提交(可多次调用,每轮提交一部分批次)。
  48. 待分批节点数:{unassigned_count}(请通过 query_global_heat_tree 查看,勿依赖本消息枚举 ID)。
  49. 剩余可新增批次数:{remaining_quota}(每日总上限 {MAX_DAILY_BATCHES},当前已存在 {planning_state["existing_groups"]} 个批次)。
  50. 每批最多 20 个节点;节点较多时分多轮 save_grade_plan,每轮关注工具返回的 unassigned_category_ids 与 remaining_batch_quota。
  51. """
  52. )
  53. except Exception:
  54. logger.exception("统筹 Agent 运行异常: biz_dt=%s", biz_dt)
  55. return
  56. summary = _summarize_agent_saves(result)
  57. if summary["save_count"] == 0:
  58. logger.warning("统筹 Agent 未成功入库任何批次: biz_dt=%s", biz_dt)
  59. return
  60. logger.info(
  61. "统筹 Agent 完成: biz_dt=%s save_count=%s persisted_groups=%s",
  62. biz_dt,
  63. summary["save_count"],
  64. summary["persisted_groups"],
  65. )
  66. last = summary["last_payload"] or {}
  67. if last.get("unassigned_category_ids"):
  68. logger.info(
  69. "统筹后仍有未分批节点: biz_dt=%s remaining=%s quota=%s",
  70. biz_dt,
  71. len(last["unassigned_category_ids"]),
  72. last.get("remaining_batch_quota"),
  73. )
  74. def orchestrate_daily_grade_plan(*, biz_dt: str) -> None:
  75. """运行统筹规划(批次由 Agent 划分,save_grade_plan 工具入库)。"""
  76. planning_state = resolve_planning_state(biz_dt)
  77. if planning_state["batch_limit_reached"]:
  78. logger.info(
  79. "跳过统筹 Agent:当天批次已达上限 biz_dt=%s existing_groups=%s limit=%s",
  80. biz_dt,
  81. planning_state["existing_groups"],
  82. MAX_DAILY_BATCHES,
  83. )
  84. return
  85. unassigned_ids = planning_state["unassigned_category_ids"]
  86. if not unassigned_ids:
  87. logger.info(
  88. "跳过统筹 Agent:当天有需求节点均已分批 biz_dt=%s total_hanging_nodes=%s",
  89. biz_dt,
  90. planning_state["total_hanging_nodes"],
  91. )
  92. return
  93. logger.info(
  94. "执行统筹 Agent: biz_dt=%s unassigned=%s remaining_quota=%s",
  95. biz_dt,
  96. len(unassigned_ids),
  97. planning_state["remaining_batch_quota"],
  98. )
  99. _orchestrate_via_agent(
  100. biz_dt,
  101. planning_state=planning_state,
  102. unassigned_count=len(unassigned_ids),
  103. )
  104. def main(biz_dt: str) -> dict[str, Any]:
  105. """手动测试:运行统筹规划并打印当天分配摘要。"""
  106. planning_before = resolve_planning_state(biz_dt)
  107. orchestrate_daily_grade_plan(biz_dt=biz_dt)
  108. planning_after = resolve_planning_state(biz_dt)
  109. with get_session() as session:
  110. repo = DemandGradePlanRepository(session)
  111. plan = repo.get_latest_plan(biz_dt)
  112. groups = repo.list_groups_by_biz_dt(biz_dt)
  113. group_status = repo.summarize(biz_dt)
  114. plan_summary = json.loads(plan.plan_json) if plan is not None else {}
  115. result = {
  116. "biz_dt": biz_dt,
  117. "planning_before": planning_before,
  118. "planning_after": planning_after,
  119. "plan_count": 1 if plan is not None else 0,
  120. "total_hanging_nodes": planning_after["total_hanging_nodes"],
  121. "group_count": len(groups),
  122. "group_status": group_status,
  123. "unassigned_category_ids": planning_after["unassigned_category_ids"],
  124. "sample_groups": [
  125. {
  126. "group_key": group.group_key,
  127. "category_ids": json.loads(group.category_ids),
  128. "status": group.status,
  129. }
  130. for group in groups[:3]
  131. ],
  132. "latest_plan_summary": plan_summary,
  133. }
  134. print(json.dumps(result, ensure_ascii=False, indent=2))
  135. return result
  136. if __name__ == "__main__":
  137. import sys
  138. main(sys.argv[1] if len(sys.argv) > 1 else "20260714")