tools.py 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322
  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 pydantic import BaseModel, Field, model_validator
  7. from find_agent_v2.qwen_video_understanding_30s import understand_candidate_video_30s_v2
  8. from find_agent_v2.providers import (
  9. fetch_details,
  10. fetch_portraits,
  11. search_internal,
  12. search_tikhub,
  13. )
  14. from find_agent_v2.service import get_find_agent_v2_service
  15. from supply_agent.tools import tool
  16. from supply_agent.tools.registry import ToolRegistry
  17. ToolFn = Callable[..., Any]
  18. class CandidateEvaluation(BaseModel):
  19. """One candidate decision. Prefer candidate_id; aweme_id is a safe fallback."""
  20. candidate_id: int | None = Field(
  21. default=None,
  22. description="数据库候选编号,即输入中的 candidate_id;不要填写 aweme_id。",
  23. )
  24. aweme_id: str | None = Field(
  25. default=None,
  26. description="抖音视频号;仅在无法填写 candidate_id 时作为兼容定位字段。",
  27. )
  28. relevance_score: float = Field(ge=0, le=1, description="R 相关性评分,0~1。")
  29. elder_score: float = Field(ge=0, le=1, description="E 中老年适配评分,0~1。")
  30. share_score: float = Field(ge=0, le=1, description="S 传播力评分,0~1。")
  31. value_score: float = Field(ge=0, le=1, description="V 综合价值评分,0~1。")
  32. decision_bucket: str = Field(description="只能是 primary 或 rejected。")
  33. decision_reason: str = Field(min_length=1, description="基于真实证据的中文判断理由。")
  34. reject_reason_code: str | None = None
  35. @model_validator(mode="after")
  36. def validate_identity_and_bucket(self):
  37. if self.candidate_id is None and not str(self.aweme_id or "").strip():
  38. raise ValueError("candidate_id 和 aweme_id 至少提供一个")
  39. if self.decision_bucket not in {"primary", "rejected"}:
  40. raise ValueError("decision_bucket 只能是 primary/rejected")
  41. return self
  42. class EvaluationBatch(BaseModel):
  43. """Structured evaluator output consumed and persisted by the host."""
  44. items: list[CandidateEvaluation] = Field(
  45. min_length=1,
  46. description="当前 Worker 分片内全部候选的评估结果,每个候选恰好一条。",
  47. )
  48. def normalize_evaluation_items(
  49. items: list[dict[str, Any] | CandidateEvaluation],
  50. *,
  51. allowed_candidates: list[dict[str, Any]],
  52. ) -> list[dict[str, Any]]:
  53. """Validate a complete worker shard and resolve aweme_id without widening access."""
  54. allowed_ids = {int(item["candidate_id"]) for item in allowed_candidates}
  55. aweme_to_id = {
  56. str(item.get("aweme_id") or ""): int(item["candidate_id"])
  57. for item in allowed_candidates
  58. }
  59. normalized: list[dict[str, Any]] = []
  60. for raw in items:
  61. value = raw if isinstance(raw, CandidateEvaluation) else CandidateEvaluation.model_validate(raw)
  62. candidate_id = value.candidate_id
  63. if candidate_id is None:
  64. candidate_id = aweme_to_id.get(str(value.aweme_id or ""))
  65. if candidate_id not in allowed_ids:
  66. raise ValueError(
  67. "评估项不属于当前 Worker 分片;"
  68. f"candidate_id={candidate_id}, aweme_id={value.aweme_id}, "
  69. f"allowed_candidate_ids={sorted(allowed_ids)}"
  70. )
  71. payload = value.model_dump(exclude_none=True)
  72. payload["candidate_id"] = candidate_id
  73. payload.pop("aweme_id", None)
  74. normalized.append(payload)
  75. requested = [int(item["candidate_id"]) for item in normalized]
  76. if len(requested) != len(set(requested)):
  77. raise ValueError("同一 candidate_id 不能重复评估")
  78. missing = allowed_ids - set(requested)
  79. if missing:
  80. raise ValueError(f"必须一次评估完整分片,缺少 candidate_ids={sorted(missing)}")
  81. return normalized
  82. def bound_candidate_tools(
  83. functions: tuple[ToolFn, ...], *, run_id: str, candidate_ids: list[int],
  84. ) -> tuple[ToolFn, ...]:
  85. """Physically constrain worker tools to one run and one candidate shard."""
  86. allowed = {int(value) for value in candidate_ids}
  87. output: list[ToolFn] = []
  88. for original in functions:
  89. name = str(getattr(original, "_tool_name", original.__name__))
  90. if name == "query_pending_candidates_v2":
  91. def make_query(fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  92. @tool(name=fn_name, description="仅查询当前 Worker 分片中的待评估候选。")
  93. def bound_query(run_id: str, limit: int = 100) -> str:
  94. if run_id != worker_run_id:
  95. return json.dumps({"error": "run_id 不属于当前 Worker"}, ensure_ascii=False)
  96. state = get_find_agent_v2_service().get_full_state(
  97. run_id, limit=limit, pending_only=True,
  98. )
  99. state["candidates"] = [
  100. item for item in state["candidates"]
  101. if int(item["candidate_id"]) in worker_allowed
  102. ]
  103. return json.dumps(state, ensure_ascii=False, default=str)
  104. return bound_query
  105. output.append(make_query(name, run_id, allowed))
  106. continue
  107. def make_async(fn: ToolFn, fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  108. @tool(name=fn_name, description=getattr(fn, "_tool_description", ""))
  109. async def bound_async(run_id: str, candidate_ids: list[int]) -> str:
  110. if run_id != worker_run_id:
  111. return json.dumps({"error": "run_id 不属于当前 Worker"}, ensure_ascii=False)
  112. requested = [int(value) for value in candidate_ids]
  113. if not requested or not set(requested) <= worker_allowed:
  114. return json.dumps({
  115. "error": "candidate_ids 超出当前 Worker 分片",
  116. "allowed_candidate_ids": sorted(worker_allowed),
  117. }, ensure_ascii=False)
  118. return await fn(run_id=run_id, candidate_ids=requested)
  119. return bound_async
  120. def make_video_understanding(
  121. fn: ToolFn, fn_name: str, worker_run_id: str, worker_allowed: set[int],
  122. ):
  123. @tool(name=fn_name, description=getattr(fn, "_tool_description", ""))
  124. async def bound_video_understanding(
  125. run_id: str, candidate_id: int, prompt: str,
  126. ) -> str:
  127. if run_id != worker_run_id or int(candidate_id) not in worker_allowed:
  128. return json.dumps({
  129. "error": "candidate_id 或 run_id 超出当前评估 Worker 分片",
  130. "allowed_candidate_ids": sorted(worker_allowed),
  131. }, ensure_ascii=False)
  132. return await fn(run_id=run_id, candidate_id=int(candidate_id), prompt=prompt)
  133. return bound_video_understanding
  134. def make_sync(fn: ToolFn, fn_name: str, worker_run_id: str, worker_allowed: set[int]):
  135. @tool(name=fn_name, description=getattr(fn, "_tool_description", ""))
  136. def bound_sync(run_id: str, items: list[CandidateEvaluation]) -> str:
  137. if run_id != worker_run_id:
  138. raise ValueError("run_id 不属于当前 Worker")
  139. state = get_find_agent_v2_service().get_full_state(
  140. run_id, pending_only=True,
  141. )
  142. allowed_candidates = [
  143. item for item in state.get("candidates", [])
  144. if int(item["candidate_id"]) in worker_allowed
  145. ]
  146. normalized = normalize_evaluation_items(
  147. items, allowed_candidates=allowed_candidates,
  148. )
  149. return fn(run_id=run_id, items=normalized)
  150. return bound_sync
  151. wrapper: ToolFn
  152. if name == "understand_candidate_video_30s_v2":
  153. wrapper = make_video_understanding(original, name, run_id, allowed)
  154. elif name in {"fetch_candidate_details_v2", "fetch_candidate_portraits_v2"}:
  155. wrapper = make_async(original, name, run_id, allowed)
  156. elif name == "evaluate_candidates_v2":
  157. wrapper = make_sync(original, name, run_id, allowed)
  158. else:
  159. wrapper = original
  160. output.append(wrapper)
  161. return tuple(output)
  162. @tool
  163. async def search_videos_v2(run_id: str, round_index: int, searches: list[dict[str, Any]]) -> str:
  164. """批量搜索并仅写入 find_agent_v2_search/candidate;searches 最多 6 项。"""
  165. if not searches or len(searches) > 6:
  166. return json.dumps({"error": "searches 必须为 1~6 项"}, ensure_ascii=False)
  167. service = get_find_agent_v2_service()
  168. outputs: list[dict[str, Any]] = []
  169. for raw_task in searches:
  170. keyword = str(raw_task.get("keyword") or "").strip()
  171. reason = str(raw_task.get("query_reason") or "").strip()
  172. provider = str(raw_task.get("provider") or "internal_keyword")
  173. if not keyword or not reason:
  174. outputs.append({"error": "keyword/query_reason 不能为空"})
  175. continue
  176. max_pages = max(1, min(int(raw_task.get("max_pages") or 1), 2))
  177. cursor: str | int = raw_task.get("cursor") or 0
  178. provider_search_id = str(raw_task.get("search_id") or "")
  179. backtrace = str(raw_task.get("backtrace") or "")
  180. for page_no in range(1, max_pages + 1):
  181. common = {
  182. "keyword": keyword,
  183. "content_type": str(raw_task.get("content_type") or "视频"),
  184. "sort_type": str(raw_task.get("sort_type") or "综合排序"),
  185. "publish_time": str(raw_task.get("publish_time") or "不限"),
  186. "min_duration_seconds": int(raw_task.get("min_duration_seconds") or 30),
  187. }
  188. if provider == "tikhub":
  189. payload = await search_tikhub(
  190. **common,
  191. cursor=int(cursor or 0),
  192. filter_duration=str(raw_task.get("filter_duration") or "不限"),
  193. search_id=provider_search_id,
  194. backtrace=backtrace,
  195. )
  196. else:
  197. provider = "internal_keyword"
  198. payload = await search_internal(**common, cursor=str(cursor or "0"))
  199. saved = service.save_search(
  200. run_id=run_id,
  201. round_index=int(round_index),
  202. keyword=keyword,
  203. query_reason=reason,
  204. source_type=str(raw_task.get("source_type") or "mixed"),
  205. provider=provider,
  206. cursor=str(cursor),
  207. page_no=page_no,
  208. payload=payload,
  209. )
  210. outputs.append({
  211. "keyword": keyword,
  212. "provider": provider,
  213. "page_no": page_no,
  214. "error": payload.get("error"),
  215. "has_more": bool(payload.get("has_more")),
  216. "next_cursor": payload.get("next_cursor"),
  217. **saved,
  218. })
  219. if payload.get("error") or not payload.get("has_more"):
  220. break
  221. cursor = payload.get("next_cursor") or cursor
  222. provider_search_id = str(payload.get("search_id") or provider_search_id)
  223. backtrace = str(payload.get("backtrace") or backtrace)
  224. return json.dumps({"run_id": run_id, "searches": outputs}, ensure_ascii=False)
  225. @tool
  226. async def fetch_candidate_details_v2(run_id: str, candidate_ids: list[int]) -> str:
  227. """批量获取候选详情并仅写入 find_agent_v2_candidate/evidence;最多 8 项。"""
  228. service = get_find_agent_v2_service()
  229. candidates = service.candidate_inputs(run_id, candidate_ids[:8])
  230. payload = await fetch_details([item["aweme_id"] for item in candidates])
  231. service.save_details(run_id, list(payload.get("details") or []), list(payload.get("errors") or []))
  232. return json.dumps({
  233. "run_id": run_id,
  234. "success_count": int(payload.get("success_count") or len(payload.get("details") or [])),
  235. "failed_count": int(payload.get("failed_count") or len(payload.get("errors") or [])),
  236. "errors": payload.get("errors") or [],
  237. }, ensure_ascii=False)
  238. @tool
  239. async def fetch_candidate_portraits_v2(run_id: str, candidate_ids: list[int]) -> str:
  240. """批量获取候选双侧年龄画像并仅写入 find_agent_v2_candidate/evidence。"""
  241. service = get_find_agent_v2_service()
  242. candidates = service.candidate_inputs(run_id, candidate_ids[:8])
  243. payload = await fetch_portraits([{
  244. "aweme_id": item["aweme_id"],
  245. "author_sec_uid": item.get("author_sec_uid"),
  246. } for item in candidates])
  247. results = list(payload.get("results") or [])
  248. service.save_portraits(run_id, results)
  249. return json.dumps({"run_id": run_id, "count": len(results), "results": results}, ensure_ascii=False)
  250. @tool
  251. def evaluate_candidates_v2(run_id: str, items: list[CandidateEvaluation]) -> str:
  252. """按 candidate_id 写入完整分片的 R/E/S/V 与分池;不要用 aweme_id 代替。"""
  253. payload = [item.model_dump(exclude_none=True) if isinstance(item, CandidateEvaluation) else item for item in items]
  254. updated = get_find_agent_v2_service().evaluate(run_id, payload)
  255. return json.dumps({"run_id": run_id, "updated": updated}, ensure_ascii=False)
  256. @tool
  257. def query_find_agent_v2_state(run_id: str, limit: int = 100) -> str:
  258. """查询完全隔离的 find_agent_v2 运行、搜索与候选状态。"""
  259. state = get_find_agent_v2_service().get_full_state(run_id, limit=limit)
  260. return json.dumps(state, ensure_ascii=False, default=str)
  261. @tool
  262. def query_pending_candidates_v2(run_id: str, limit: int = 100) -> str:
  263. """仅查询当前 run 尚未分池的 pending_evaluation 候选。"""
  264. state = get_find_agent_v2_service().get_full_state(
  265. run_id, limit=limit, pending_only=True,
  266. )
  267. return json.dumps(state, ensure_ascii=False, default=str)
  268. SEARCH_TOOLS: tuple[ToolFn, ...] = (search_videos_v2, query_find_agent_v2_state)
  269. EVIDENCE_TOOLS: tuple[ToolFn, ...] = (
  270. fetch_candidate_details_v2,
  271. fetch_candidate_portraits_v2,
  272. query_pending_candidates_v2,
  273. )
  274. EVALUATION_TOOLS: tuple[ToolFn, ...] = (
  275. understand_candidate_video_30s_v2,
  276. evaluate_candidates_v2,
  277. query_pending_candidates_v2,
  278. )
  279. REPORT_TOOLS: tuple[ToolFn, ...] = (query_find_agent_v2_state,)
  280. def build_tool_registry(functions: Iterable[ToolFn]) -> ToolRegistry:
  281. return ToolRegistry().from_decorated(*tuple(functions))