search_persistence.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301
  1. """搜索结果的内部持久化流程,不暴露为 Agent 工具。"""
  2. from __future__ import annotations
  3. import hashlib
  4. import json
  5. import logging
  6. from typing import Any
  7. from supply_infra.video_discovery_gates import (
  8. normalize_duration_seconds,
  9. parse_datetime_value,
  10. )
  11. from supply_infra.services.video_discovery_service import (
  12. RunNotFoundError,
  13. build_evidence_values,
  14. format_db_error,
  15. get_video_discovery_service,
  16. )
  17. logger = logging.getLogger(__name__)
  18. _SOURCE_TYPES = {
  19. "demand",
  20. "seed",
  21. "point",
  22. "tag",
  23. "author",
  24. "pagination",
  25. "mixed",
  26. }
  27. _ROOT_SOURCE_TYPES = {"demand", "seed", "point", "mixed"}
  28. def _load_payload(payload_json: str) -> dict[str, Any]:
  29. try:
  30. payload = json.loads(payload_json)
  31. except (TypeError, ValueError):
  32. return {"error": "搜索接口返回了无效 JSON", "raw_result": str(payload_json)}
  33. if not isinstance(payload, dict):
  34. return {"error": "搜索接口返回值不是对象", "raw_result": payload}
  35. return payload
  36. def _clean_text(value: Any, *, max_length: int | None = None) -> str | None:
  37. if value is None:
  38. return None
  39. text = str(value).strip()
  40. if not text:
  41. return None
  42. return text[:max_length] if max_length else text
  43. def _nonnegative_int(value: Any) -> int | None:
  44. if value is None or value == "":
  45. return None
  46. try:
  47. return max(0, int(value))
  48. except (TypeError, ValueError):
  49. return None
  50. def _first_present(*values: Any) -> Any:
  51. for value in values:
  52. if value is not None and value != "":
  53. return value
  54. return None
  55. def _search_key(values: dict[str, Any]) -> str:
  56. identity = {
  57. key: values.get(key)
  58. for key in (
  59. "provider",
  60. "keyword",
  61. "content_type",
  62. "sort_type",
  63. "publish_time",
  64. "cursor",
  65. )
  66. }
  67. raw = json.dumps(identity, ensure_ascii=False, sort_keys=True)
  68. return hashlib.sha256(raw.encode("utf-8")).hexdigest()
  69. def _candidate_from_search_result(
  70. item: dict[str, Any],
  71. keyword: str,
  72. ) -> dict[str, Any] | None:
  73. aweme_id = _clean_text(
  74. item.get("aweme_id") or item.get("content_id"),
  75. max_length=64,
  76. )
  77. if not aweme_id:
  78. return None
  79. author = item.get("author") if isinstance(item.get("author"), dict) else {}
  80. stats = item.get("statistics") if isinstance(item.get("statistics"), dict) else {}
  81. topics = item.get("topics") if isinstance(item.get("topics"), list) else []
  82. duration = normalize_duration_seconds(
  83. item.get("duration_seconds")
  84. if item.get("duration_seconds") not in (None, "")
  85. else item.get("duration_ms"),
  86. unit=(
  87. "seconds"
  88. if item.get("duration_seconds") not in (None, "")
  89. else "milliseconds"
  90. ),
  91. )
  92. publish_at = parse_datetime_value(
  93. item.get("publish_at")
  94. or item.get("create_time")
  95. or item.get("create_timestamp")
  96. or item.get("publish_timestamp")
  97. )
  98. return {
  99. "aweme_id": aweme_id,
  100. "title": _clean_text(item.get("desc") or item.get("title"), max_length=512),
  101. "content_link": _clean_text(
  102. item.get("url") or item.get("content_link"),
  103. max_length=1024,
  104. ),
  105. "author_name": _clean_text(
  106. author.get("nickname") or item.get("author_name"),
  107. max_length=256,
  108. ),
  109. "author_sec_uid": _clean_text(
  110. author.get("sec_uid") or item.get("author_sec_uid"),
  111. max_length=256,
  112. ),
  113. "like_count": _nonnegative_int(
  114. _first_present(stats.get("digg_count"), item.get("like_count"))
  115. ),
  116. "comment_count": _nonnegative_int(
  117. _first_present(stats.get("comment_count"), item.get("comment_count"))
  118. ),
  119. "share_count": _nonnegative_int(
  120. _first_present(stats.get("share_count"), item.get("share_count"))
  121. ),
  122. "collect_count": _nonnegative_int(
  123. _first_present(stats.get("collect_count"), item.get("collect_count"))
  124. ),
  125. "play_count": _nonnegative_int(
  126. _first_present(stats.get("play_count"), item.get("play_count"))
  127. ),
  128. "duration_seconds": duration,
  129. "publish_at": (
  130. publish_at.replace(tzinfo=None) if publish_at is not None else None
  131. ),
  132. "tags_json": (
  133. json.dumps(topics, ensure_ascii=False) if topics else None
  134. ),
  135. "_source_keyword": keyword,
  136. }
  137. def persist_search_payload(
  138. payload_json: str,
  139. *,
  140. run_id: str,
  141. keyword: str,
  142. query_reason: str,
  143. source_type: str,
  144. cursor: str,
  145. page_no: int,
  146. provider: str,
  147. source_value: str | None = None,
  148. parent_search_id: int | None = None,
  149. content_type: str = "视频",
  150. sort_type: str = "综合排序",
  151. publish_time: str = "不限",
  152. provider_state: dict[str, Any] | None = None,
  153. ) -> str:
  154. """新增搜索记录和本页全部候选,并把数据库 ID 拼回搜索结果。"""
  155. payload = _load_payload(payload_json)
  156. provider_raw_response = payload.get("_raw_response", payload)
  157. run_text = _clean_text(run_id, max_length=64)
  158. keyword_text = _clean_text(keyword, max_length=256)
  159. reason_text = _clean_text(query_reason)
  160. if not run_text or not keyword_text or not reason_text:
  161. payload["error"] = "run_id、keyword、query_reason 不能为空"
  162. payload["input_error"] = True
  163. return json.dumps(payload, ensure_ascii=False, default=str)
  164. if source_type not in _SOURCE_TYPES:
  165. payload["error"] = f"source_type 必须是: {sorted(_SOURCE_TYPES)}"
  166. payload["input_error"] = True
  167. return json.dumps(payload, ensure_ascii=False, default=str)
  168. raw_results = payload.get("search_results")
  169. results = raw_results if isinstance(raw_results, list) else []
  170. candidate_rows = [
  171. row
  172. for item in results
  173. if isinstance(item, dict)
  174. if (row := _candidate_from_search_result(dict(item), keyword_text)) is not None
  175. ]
  176. normalized_page_no = max(1, int(page_no))
  177. normalized_source_type = (
  178. "pagination" if normalized_page_no > 1 else source_type
  179. )
  180. normalized_parent_search_id = (
  181. None
  182. if normalized_source_type in _ROOT_SOURCE_TYPES
  183. else parent_search_id
  184. )
  185. merged_provider_state = dict(provider_state or {})
  186. for key in ("search_id", "backtrace"):
  187. value = payload.get(key)
  188. if value not in (None, ""):
  189. merged_provider_state[key] = value
  190. provider_search_id = payload.get("search_id")
  191. search_values: dict[str, Any] = {
  192. "run_id": run_text,
  193. "keyword": keyword_text,
  194. "query_reason": reason_text,
  195. "source_type": normalized_source_type,
  196. "source_value": _clean_text(source_value),
  197. "parent_search_id": normalized_parent_search_id,
  198. "provider": _clean_text(provider, max_length=32) or "internal_keyword",
  199. "provider_state_json": (
  200. json.dumps(merged_provider_state, ensure_ascii=False)
  201. if merged_provider_state
  202. else None
  203. ),
  204. "content_type": _clean_text(content_type, max_length=16) or "视频",
  205. "sort_type": _clean_text(sort_type, max_length=32) or "综合排序",
  206. "publish_time": _clean_text(publish_time, max_length=32) or "不限",
  207. "cursor": _clean_text(cursor, max_length=128) or "0",
  208. "page_no": normalized_page_no,
  209. "results_count": len(results),
  210. "new_candidate_count": 0,
  211. "has_more": int(bool(payload.get("has_more"))),
  212. "next_cursor": (
  213. _clean_text(payload.get("next_cursor"), max_length=128)
  214. ),
  215. "result_ids_json": None,
  216. "status": "failed" if payload.get("error") else "success",
  217. "error_message": _clean_text(payload.get("error")),
  218. }
  219. search_values["search_key"] = _search_key(search_values)
  220. evidence_values = build_evidence_values(
  221. run_id=run_text,
  222. evidence_type="search_page",
  223. provider=str(search_values["provider"]),
  224. subject_key=f"search:{search_values['search_key']}:{normalized_page_no}",
  225. request={
  226. "keyword": keyword_text,
  227. "content_type": search_values["content_type"],
  228. "sort_type": search_values["sort_type"],
  229. "publish_time": search_values["publish_time"],
  230. "cursor": search_values["cursor"],
  231. "page_no": normalized_page_no,
  232. },
  233. raw_response=provider_raw_response,
  234. normalized={"candidate_count": len(candidate_rows)},
  235. fetch_status="failed" if payload.get("error") else "success",
  236. error_message=_clean_text(payload.get("error")),
  237. )
  238. try:
  239. saved = get_video_discovery_service().save_search_page(
  240. run_text,
  241. search_values,
  242. candidate_rows,
  243. evidence_values,
  244. )
  245. except RunNotFoundError as exc:
  246. payload["error"] = str(exc)
  247. payload["input_error"] = True
  248. return json.dumps(payload, ensure_ascii=False, default=str)
  249. except Exception as exc:
  250. logger.error("persist search payload failed: %s", exc, exc_info=True)
  251. payload["error"] = format_db_error(exc)
  252. return json.dumps(payload, ensure_ascii=False, default=str)
  253. payload.pop("search_results", None)
  254. payload.pop("user_videos", None)
  255. payload.pop("_raw_response", None)
  256. if provider_search_id not in (None, ""):
  257. payload["provider_search_id"] = provider_search_id
  258. saved_candidates = saved.get("candidates") or []
  259. saved["candidates"] = [
  260. {
  261. "candidate_id": item.get("candidate_id"),
  262. "search_id": item.get("search_id"),
  263. "aweme_id": item.get("aweme_id"),
  264. "title": item.get("title"),
  265. "author_name": item.get("author_name"),
  266. "author_sec_uid": item.get("author_sec_uid"),
  267. "publish_at": item.get("publish_at"),
  268. "duration_seconds": item.get("duration_seconds"),
  269. "share_count": item.get("share_count"),
  270. "gate_status": item.get("gate_status"),
  271. }
  272. for item in saved_candidates
  273. ]
  274. payload.update(saved)
  275. payload["persisted"] = True
  276. return json.dumps(payload, ensure_ascii=False, default=str)