"""运行统筹规划 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")