tools.py 9.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222
  1. """Stage tools that persist exclusively to ``find_agent_v2_*`` tables."""
  2. from __future__ import annotations
  3. import json
  4. from collections.abc import Callable, Iterable
  5. from typing import Any
  6. from find_agent_v2.providers import (
  7. fetch_details,
  8. fetch_portraits,
  9. search_internal,
  10. search_tikhub,
  11. )
  12. from find_agent_v2.service import get_find_agent_v2_service
  13. from supply_agent.tools import tool
  14. from supply_agent.tools.registry import ToolRegistry
  15. ToolFn = Callable[..., Any]
  16. def bound_candidate_tools(
  17. functions: tuple[ToolFn, ...], *, run_id: str, candidate_ids: list[int],
  18. ) -> tuple[ToolFn, ...]:
  19. """Physically constrain worker tools to one run and one candidate shard."""
  20. allowed = {int(value) for value in candidate_ids}
  21. output: list[ToolFn] = []
  22. for original in functions:
  23. name = str(getattr(original, "_tool_name", original.__name__))
  24. if name == "query_pending_candidates_v2":
  25. def make_query(fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  26. @tool(name=fn_name, description="仅查询当前 Worker 分片中的待评估候选。")
  27. def bound_query(run_id: str, limit: int = 100) -> str:
  28. if run_id != worker_run_id:
  29. return json.dumps({"error": "run_id 不属于当前 Worker"}, ensure_ascii=False)
  30. state = get_find_agent_v2_service().get_full_state(
  31. run_id, limit=limit, pending_only=True,
  32. )
  33. state["candidates"] = [
  34. item for item in state["candidates"]
  35. if int(item["candidate_id"]) in worker_allowed
  36. ]
  37. return json.dumps(state, ensure_ascii=False, default=str)
  38. return bound_query
  39. output.append(make_query(name, run_id, allowed))
  40. continue
  41. def make_async(fn: ToolFn, fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  42. @tool(name=fn_name, description=getattr(fn, "_tool_description", ""))
  43. async def bound_async(run_id: str, candidate_ids: list[int]) -> str:
  44. if run_id != worker_run_id:
  45. return json.dumps({"error": "run_id 不属于当前 Worker"}, ensure_ascii=False)
  46. requested = [int(value) for value in candidate_ids]
  47. if not requested or not set(requested) <= worker_allowed:
  48. return json.dumps({
  49. "error": "candidate_ids 超出当前 Worker 分片",
  50. "allowed_candidate_ids": sorted(worker_allowed),
  51. }, ensure_ascii=False)
  52. return await fn(run_id=run_id, candidate_ids=requested)
  53. return bound_async
  54. def make_sync(fn: ToolFn, fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  55. @tool(name=fn_name, description=getattr(fn, "_tool_description", ""))
  56. def bound_sync(run_id: str, items: list[dict[str, Any]]) -> str:
  57. if run_id != worker_run_id:
  58. return json.dumps({"error": "run_id 不属于当前 Worker"}, ensure_ascii=False)
  59. requested = {int(item.get("candidate_id") or 0) for item in items}
  60. if not requested or not requested <= worker_allowed:
  61. return json.dumps({
  62. "error": "items 超出当前 Worker 分片",
  63. "allowed_candidate_ids": sorted(worker_allowed),
  64. }, ensure_ascii=False)
  65. return fn(run_id=run_id, items=items)
  66. return bound_sync
  67. wrapper: ToolFn
  68. if name in {"fetch_candidate_details_v2", "fetch_candidate_portraits_v2"}:
  69. wrapper = make_async(original, name, run_id, allowed)
  70. elif name == "evaluate_candidates_v2":
  71. wrapper = make_sync(original, name, run_id, allowed)
  72. else:
  73. wrapper = original
  74. output.append(wrapper)
  75. return tuple(output)
  76. @tool
  77. async def search_videos_v2(run_id: str, round_index: int, searches: list[dict[str, Any]]) -> str:
  78. """批量搜索并仅写入 find_agent_v2_search/candidate;searches 最多 6 项。"""
  79. if not searches or len(searches) > 6:
  80. return json.dumps({"error": "searches 必须为 1~6 项"}, ensure_ascii=False)
  81. service = get_find_agent_v2_service()
  82. outputs: list[dict[str, Any]] = []
  83. for raw_task in searches:
  84. keyword = str(raw_task.get("keyword") or "").strip()
  85. reason = str(raw_task.get("query_reason") or "").strip()
  86. provider = str(raw_task.get("provider") or "internal_keyword")
  87. if not keyword or not reason:
  88. outputs.append({"error": "keyword/query_reason 不能为空"})
  89. continue
  90. max_pages = max(1, min(int(raw_task.get("max_pages") or 1), 2))
  91. cursor: str | int = raw_task.get("cursor") or 0
  92. provider_search_id = str(raw_task.get("search_id") or "")
  93. backtrace = str(raw_task.get("backtrace") or "")
  94. for page_no in range(1, max_pages + 1):
  95. common = {
  96. "keyword": keyword,
  97. "content_type": str(raw_task.get("content_type") or "视频"),
  98. "sort_type": str(raw_task.get("sort_type") or "综合排序"),
  99. "publish_time": str(raw_task.get("publish_time") or "不限"),
  100. "min_duration_seconds": int(raw_task.get("min_duration_seconds") or 30),
  101. }
  102. if provider == "tikhub":
  103. payload = await search_tikhub(
  104. **common,
  105. cursor=int(cursor or 0),
  106. filter_duration=str(raw_task.get("filter_duration") or "不限"),
  107. search_id=provider_search_id,
  108. backtrace=backtrace,
  109. )
  110. else:
  111. provider = "internal_keyword"
  112. payload = await search_internal(**common, cursor=str(cursor or "0"))
  113. saved = service.save_search(
  114. run_id=run_id,
  115. round_index=int(round_index),
  116. keyword=keyword,
  117. query_reason=reason,
  118. source_type=str(raw_task.get("source_type") or "mixed"),
  119. provider=provider,
  120. cursor=str(cursor),
  121. page_no=page_no,
  122. payload=payload,
  123. )
  124. outputs.append({
  125. "keyword": keyword,
  126. "provider": provider,
  127. "page_no": page_no,
  128. "error": payload.get("error"),
  129. "has_more": bool(payload.get("has_more")),
  130. "next_cursor": payload.get("next_cursor"),
  131. **saved,
  132. })
  133. if payload.get("error") or not payload.get("has_more"):
  134. break
  135. cursor = payload.get("next_cursor") or cursor
  136. provider_search_id = str(payload.get("search_id") or provider_search_id)
  137. backtrace = str(payload.get("backtrace") or backtrace)
  138. return json.dumps({"run_id": run_id, "searches": outputs}, ensure_ascii=False)
  139. @tool
  140. async def fetch_candidate_details_v2(run_id: str, candidate_ids: list[int]) -> str:
  141. """批量获取候选详情并仅写入 find_agent_v2_candidate/evidence;最多 8 项。"""
  142. service = get_find_agent_v2_service()
  143. candidates = service.candidate_inputs(run_id, candidate_ids[:8])
  144. payload = await fetch_details([item["aweme_id"] for item in candidates])
  145. service.save_details(run_id, list(payload.get("details") or []), list(payload.get("errors") or []))
  146. return json.dumps({
  147. "run_id": run_id,
  148. "success_count": int(payload.get("success_count") or len(payload.get("details") or [])),
  149. "failed_count": int(payload.get("failed_count") or len(payload.get("errors") or [])),
  150. "errors": payload.get("errors") or [],
  151. }, ensure_ascii=False)
  152. @tool
  153. async def fetch_candidate_portraits_v2(run_id: str, candidate_ids: list[int]) -> str:
  154. """批量获取候选双侧年龄画像并仅写入 find_agent_v2_candidate/evidence。"""
  155. service = get_find_agent_v2_service()
  156. candidates = service.candidate_inputs(run_id, candidate_ids[:8])
  157. payload = await fetch_portraits([{
  158. "aweme_id": item["aweme_id"],
  159. "author_sec_uid": item.get("author_sec_uid"),
  160. } for item in candidates])
  161. results = list(payload.get("results") or [])
  162. service.save_portraits(run_id, results)
  163. return json.dumps({"run_id": run_id, "count": len(results), "results": results}, ensure_ascii=False)
  164. @tool
  165. def evaluate_candidates_v2(run_id: str, items: list[dict[str, Any]]) -> str:
  166. """写入 R/E/S/V 与 primary/rejected;程序会对 primary 强制执行 P0 门禁。"""
  167. updated = get_find_agent_v2_service().evaluate(run_id, items)
  168. return json.dumps({"run_id": run_id, "updated": updated}, ensure_ascii=False)
  169. @tool
  170. def query_find_agent_v2_state(run_id: str, limit: int = 100) -> str:
  171. """查询完全隔离的 find_agent_v2 运行、搜索与候选状态。"""
  172. state = get_find_agent_v2_service().get_full_state(run_id, limit=limit)
  173. return json.dumps(state, ensure_ascii=False, default=str)
  174. @tool
  175. def query_pending_candidates_v2(run_id: str, limit: int = 100) -> str:
  176. """仅查询当前 run 尚未分池的 pending_evaluation 候选。"""
  177. state = get_find_agent_v2_service().get_full_state(
  178. run_id, limit=limit, pending_only=True,
  179. )
  180. return json.dumps(state, ensure_ascii=False, default=str)
  181. SEARCH_TOOLS: tuple[ToolFn, ...] = (search_videos_v2, query_find_agent_v2_state)
  182. EVIDENCE_TOOLS: tuple[ToolFn, ...] = (
  183. fetch_candidate_details_v2,
  184. fetch_candidate_portraits_v2,
  185. query_pending_candidates_v2,
  186. )
  187. EVALUATION_TOOLS: tuple[ToolFn, ...] = (
  188. evaluate_candidates_v2,
  189. query_pending_candidates_v2,
  190. )
  191. REPORT_TOOLS: tuple[ToolFn, ...] = (query_find_agent_v2_state,)
  192. def build_tool_registry(functions: Iterable[ToolFn]) -> ToolRegistry:
  193. return ToolRegistry().from_decorated(*tuple(functions))