agent.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189
  1. """
  2. find_agent 工厂 — 组装 Agent 实例。
  3. 每个业务 Agent 都应提供 create_xxx_agent() 工厂函数,
  4. 统一注册本 Agent 的工具 + 共享基础设施工具。
  5. """
  6. from __future__ import annotations
  7. import json
  8. import re
  9. from pathlib import Path
  10. from supply_agent import Agent
  11. from supply_agent.config import Settings
  12. from supply_agent.types import Message, Role
  13. from agents.find_agent.tools import register_all_tools
  14. _PROMPT_PATH = Path(__file__).parent / "prompt" / "system_prompt.md"
  15. FIND_AGENT_SYSTEM_PROMPT = _PROMPT_PATH.read_text(encoding="utf-8")
  16. _FAILURE_REPORT_PREFIX = "任务未完成(工具故障)"
  17. _VIDEO_ID_PATTERN = re.compile(r"(?<!\d)\d{15,22}(?!\d)")
  18. def _report_section(
  19. content: str,
  20. start_marker: str,
  21. end_markers: tuple[str, ...],
  22. ) -> str | None:
  23. start = content.find(start_marker)
  24. if start < 0:
  25. return None
  26. ends = [
  27. position
  28. for marker in end_markers
  29. if (position := content.find(marker, start + len(start_marker))) >= 0
  30. ]
  31. end = min(ends) if ends else len(content)
  32. return content[start:end]
  33. def _validate_report_primary_section(
  34. content: str,
  35. state: dict[str, object],
  36. ) -> str | None:
  37. """Only validate aweme_ids listed in the primary report section."""
  38. candidates = state.get("candidates")
  39. if not isinstance(candidates, list):
  40. return "最终状态缺少 candidates,无法校验报告分池"
  41. bucket_by_id = {
  42. str(item.get("aweme_id")): str(item.get("decision_bucket"))
  43. for item in candidates
  44. if isinstance(item, dict) and item.get("aweme_id")
  45. }
  46. primary_section = _report_section(
  47. content,
  48. "主推荐",
  49. ("淘汰候选", "搜索树", "缺失数据", "总结"),
  50. )
  51. if primary_section is None:
  52. return "最终报告必须包含“主推荐”段"
  53. primary_ids = set(_VIDEO_ID_PATTERN.findall(primary_section))
  54. if not primary_ids:
  55. return "主推荐段必须输出至少一个 aweme_id"
  56. wrong_primary = sorted(
  57. video_id
  58. for video_id in primary_ids
  59. if bucket_by_id.get(video_id) != "primary"
  60. )
  61. if wrong_primary:
  62. return (
  63. "主推荐段包含非 primary 候选: "
  64. + ", ".join(wrong_primary[:5])
  65. + "。必须按数据库 decision_bucket 输出"
  66. )
  67. return None
  68. def _primary_count_from_state(state: dict[str, object]) -> int:
  69. run = state.get("run")
  70. if isinstance(run, dict):
  71. return int(run.get("primary_count") or 0)
  72. return 0
  73. def _successful_tool_events(
  74. messages: list[Message],
  75. ) -> list[tuple[int, str, dict[str, object]]]:
  76. events: list[tuple[int, str, dict[str, object]]] = []
  77. for index, message in enumerate(messages):
  78. if message.role != Role.TOOL or not message.name or not message.content:
  79. continue
  80. try:
  81. payload = json.loads(message.content)
  82. except (TypeError, json.JSONDecodeError):
  83. continue
  84. if not isinstance(payload, dict) or payload.get("error"):
  85. continue
  86. events.append((index, message.name, payload))
  87. return events
  88. def _last_final_assistant_message(messages: list[Message]) -> Message | None:
  89. return next(
  90. (
  91. message
  92. for message in reversed(messages)
  93. if message.role == Role.ASSISTANT and not message.tool_calls
  94. ),
  95. None,
  96. )
  97. def find_agent_completion_guard(messages: list[Message]) -> str | None:
  98. """返回完成提示;默认仅提示、不阻断结束。"""
  99. last_assistant = _last_final_assistant_message(messages)
  100. content = (last_assistant.content or "").strip() if last_assistant else ""
  101. if content.startswith(_FAILURE_REPORT_PREFIX):
  102. return None
  103. events = _successful_tool_events(messages)
  104. if not any(name == "create_video_discovery_run" for _, name, _ in events):
  105. return "尚未成功创建视频发现运行"
  106. evaluation_events = [
  107. (index, payload)
  108. for index, name, payload in events
  109. if name == "batch_save_video_candidate_evaluations"
  110. ]
  111. if not evaluation_events:
  112. return "尚未保存候选评估"
  113. evaluation_index, evaluation = evaluation_events[-1]
  114. if evaluation.get("status") != "finished":
  115. return "最后一次候选保存尚未把运行状态设置为 finished"
  116. state_events = [
  117. (index, payload)
  118. for index, name, payload in events
  119. if name == "query_video_discovery_state"
  120. ]
  121. if not state_events:
  122. return "尚未查询最终数据库状态"
  123. state_index, state = state_events[-1]
  124. if state_index < evaluation_index:
  125. return "最终状态查询必须在候选保存为 finished 之后执行"
  126. run = state.get("run")
  127. run_state = run if isinstance(run, dict) else {}
  128. if run_state.get("status") != "finished":
  129. return "运行状态尚未持久化为 finished"
  130. if last_assistant is None or not content:
  131. return "最终回答为空"
  132. if _primary_count_from_state(state) > 0:
  133. return _validate_report_primary_section(content, state)
  134. return None
  135. def create_find_agent(
  136. settings: Settings | None = None,
  137. *,
  138. model: str | None = None,
  139. ) -> Agent:
  140. """创建 find_agent 实例,注册所有相关工具。"""
  141. agent = Agent(
  142. settings=settings,
  143. name="find_agent",
  144. model="google/gemini-3-flash-preview",
  145. system_prompt=FIND_AGENT_SYSTEM_PROMPT,
  146. max_iterations=60,
  147. temperature=0.2,
  148. completion_guard=find_agent_completion_guard,
  149. tool_repeat_requires_change={
  150. "query_video_discovery_state": {
  151. "record_video_search_page",
  152. "batch_save_video_candidate_evaluations",
  153. "audit_video_discovery_run",
  154. },
  155. },
  156. )
  157. # 本 Agent 专属工具
  158. register_all_tools(agent.tools)
  159. return agent