| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158 |
- """运行统筹规划 Agent(save_grade_plan 工具内直接入库)。"""
- from __future__ import annotations
- import json
- import logging
- from typing import Any
- from agents.demand_grade_orchestrator_agent import create_demand_grade_orchestrator_agent
- from agents.demand_grade_orchestrator_agent.common.assignment import MAX_DAILY_BATCHES, resolve_planning_state
- from supply_agent.types import Role
- from supply_infra.db.repositories.demand_grade_plan_repo import DemandGradePlanRepository
- from supply_infra.db.session import get_session
- logger = logging.getLogger(__name__)
- def _summarize_agent_saves(result: Any) -> dict[str, Any]:
- """统计 Agent 会话中 save_grade_plan 成功入库的次数与批次数。"""
- save_count = 0
- persisted_groups = 0
- last_payload: dict[str, Any] | None = None
- for message in result.messages:
- if message.role != Role.TOOL or message.name != "save_grade_plan":
- continue
- try:
- payload = json.loads(message.content or "")
- except json.JSONDecodeError:
- continue
- if not isinstance(payload, dict):
- continue
- if payload.get("ok") is True and payload.get("persisted") is True:
- save_count += 1
- persisted_groups += int(payload.get("persisted_group_count") or 0)
- last_payload = payload
- return {
- "save_count": save_count,
- "persisted_groups": persisted_groups,
- "last_payload": last_payload,
- }
- def _orchestrate_via_agent(
- biz_dt: str,
- *,
- planning_state: dict[str, Any],
- unassigned_count: int,
- ) -> None:
- remaining_quota = int(planning_state["remaining_batch_quota"])
- agent = create_demand_grade_orchestrator_agent()
- try:
- result = agent.run(
- f"""请为业务日 {biz_dt} 制定全局树需求分级计划。
- 必须从 query_global_heat_tree 开始,经过至少一次 query_heat_node_group 下钻,
- 由你自行划分批次后调用 save_grade_plan 提交(可多次调用,每轮提交一部分批次)。
- 待分批节点数:{unassigned_count}(请通过 query_global_heat_tree 查看,勿依赖本消息枚举 ID)。
- 剩余可新增批次数:{remaining_quota}(每日总上限 {MAX_DAILY_BATCHES},当前已存在 {planning_state["existing_groups"]} 个批次)。
- 每批最多 20 个节点;节点较多时分多轮 save_grade_plan,每轮关注工具返回的 unassigned_category_ids 与 remaining_batch_quota。
- """
- )
- except Exception:
- logger.exception("统筹 Agent 运行异常: biz_dt=%s", biz_dt)
- return
- summary = _summarize_agent_saves(result)
- if summary["save_count"] == 0:
- logger.warning("统筹 Agent 未成功入库任何批次: biz_dt=%s", biz_dt)
- return
- logger.info(
- "统筹 Agent 完成: biz_dt=%s save_count=%s persisted_groups=%s",
- biz_dt,
- summary["save_count"],
- summary["persisted_groups"],
- )
- last = summary["last_payload"] or {}
- if last.get("unassigned_category_ids"):
- logger.info(
- "统筹后仍有未分批节点: biz_dt=%s remaining=%s quota=%s",
- biz_dt,
- len(last["unassigned_category_ids"]),
- last.get("remaining_batch_quota"),
- )
- def orchestrate_daily_grade_plan(*, biz_dt: str) -> None:
- """运行统筹规划(批次由 Agent 划分,save_grade_plan 工具入库)。"""
- planning_state = resolve_planning_state(biz_dt)
- if planning_state["batch_limit_reached"]:
- logger.info(
- "跳过统筹 Agent:当天批次已达上限 biz_dt=%s existing_groups=%s limit=%s",
- biz_dt,
- planning_state["existing_groups"],
- MAX_DAILY_BATCHES,
- )
- return
- unassigned_ids = planning_state["unassigned_category_ids"]
- if not unassigned_ids:
- logger.info(
- "跳过统筹 Agent:当天有需求节点均已分批 biz_dt=%s total_hanging_nodes=%s",
- biz_dt,
- planning_state["total_hanging_nodes"],
- )
- return
- logger.info(
- "执行统筹 Agent: biz_dt=%s unassigned=%s remaining_quota=%s",
- biz_dt,
- len(unassigned_ids),
- planning_state["remaining_batch_quota"],
- )
- _orchestrate_via_agent(
- biz_dt,
- planning_state=planning_state,
- unassigned_count=len(unassigned_ids),
- )
- def main(biz_dt: str) -> dict[str, Any]:
- """手动测试:运行统筹规划并打印当天分配摘要。"""
- planning_before = resolve_planning_state(biz_dt)
- orchestrate_daily_grade_plan(biz_dt=biz_dt)
- planning_after = resolve_planning_state(biz_dt)
- with get_session() as session:
- repo = DemandGradePlanRepository(session)
- plan = repo.get_latest_plan(biz_dt)
- groups = repo.list_groups_by_biz_dt(biz_dt)
- group_status = repo.summarize(biz_dt)
- plan_summary = json.loads(plan.plan_json) if plan is not None else {}
- result = {
- "biz_dt": biz_dt,
- "planning_before": planning_before,
- "planning_after": planning_after,
- "plan_count": 1 if plan is not None else 0,
- "total_hanging_nodes": planning_after["total_hanging_nodes"],
- "group_count": len(groups),
- "group_status": group_status,
- "unassigned_category_ids": planning_after["unassigned_category_ids"],
- "sample_groups": [
- {
- "group_key": group.group_key,
- "category_ids": json.loads(group.category_ids),
- "status": group.status,
- }
- for group in groups[:3]
- ],
- "latest_plan_summary": plan_summary,
- }
- print(json.dumps(result, ensure_ascii=False, indent=2))
- return result
- if __name__ == "__main__":
- import sys
- main(sys.argv[1] if len(sys.argv) > 1 else "20260714")
|