search_persistence.py 7.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232
  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.services.video_discovery_service import (
  8. RunNotFoundError,
  9. format_db_error,
  10. get_video_discovery_service,
  11. )
  12. logger = logging.getLogger(__name__)
  13. _SOURCE_TYPES = {
  14. "demand",
  15. "seed",
  16. "point",
  17. "tag",
  18. "author",
  19. "pagination",
  20. "mixed",
  21. }
  22. _ROOT_SOURCE_TYPES = {"demand", "seed", "point", "mixed"}
  23. def _load_payload(payload_json: str) -> dict[str, Any]:
  24. try:
  25. payload = json.loads(payload_json)
  26. except (TypeError, ValueError):
  27. return {"error": "搜索接口返回了无效 JSON", "raw_result": str(payload_json)}
  28. if not isinstance(payload, dict):
  29. return {"error": "搜索接口返回值不是对象", "raw_result": payload}
  30. return payload
  31. def _clean_text(value: Any, *, max_length: int | None = None) -> str | None:
  32. if value is None:
  33. return None
  34. text = str(value).strip()
  35. if not text:
  36. return None
  37. return text[:max_length] if max_length else text
  38. def _nonnegative_int(value: Any) -> int | None:
  39. if value is None or value == "":
  40. return None
  41. try:
  42. return max(0, int(value))
  43. except (TypeError, ValueError):
  44. return None
  45. def _search_key(values: dict[str, Any]) -> str:
  46. identity = {
  47. key: values.get(key)
  48. for key in (
  49. "provider",
  50. "keyword",
  51. "content_type",
  52. "sort_type",
  53. "publish_time",
  54. "cursor",
  55. )
  56. }
  57. raw = json.dumps(identity, ensure_ascii=False, sort_keys=True)
  58. return hashlib.sha256(raw.encode("utf-8")).hexdigest()
  59. def _candidate_from_search_result(
  60. item: dict[str, Any],
  61. keyword: str,
  62. ) -> dict[str, Any] | None:
  63. aweme_id = _clean_text(
  64. item.get("aweme_id") or item.get("content_id"),
  65. max_length=64,
  66. )
  67. if not aweme_id:
  68. return None
  69. author = item.get("author") if isinstance(item.get("author"), dict) else {}
  70. stats = item.get("statistics") if isinstance(item.get("statistics"), dict) else {}
  71. topics = item.get("topics") if isinstance(item.get("topics"), list) else []
  72. return {
  73. "aweme_id": aweme_id,
  74. "title": _clean_text(item.get("desc") or item.get("title"), max_length=512),
  75. "content_link": _clean_text(
  76. item.get("url") or item.get("content_link"),
  77. max_length=1024,
  78. ),
  79. "author_name": _clean_text(
  80. author.get("nickname") or item.get("author_name"),
  81. max_length=256,
  82. ),
  83. "author_sec_uid": _clean_text(
  84. author.get("sec_uid") or item.get("author_sec_uid"),
  85. max_length=256,
  86. ),
  87. "like_count": _nonnegative_int(
  88. stats.get("digg_count") or item.get("like_count")
  89. ),
  90. "comment_count": _nonnegative_int(
  91. stats.get("comment_count") or item.get("comment_count")
  92. ),
  93. "share_count": _nonnegative_int(
  94. stats.get("share_count") or item.get("share_count")
  95. ),
  96. "collect_count": _nonnegative_int(
  97. stats.get("collect_count") or item.get("collect_count")
  98. ),
  99. "play_count": _nonnegative_int(
  100. stats.get("play_count") or item.get("play_count")
  101. ),
  102. "tags_json": (
  103. json.dumps(topics, ensure_ascii=False) if topics else None
  104. ),
  105. "_source_keyword": keyword,
  106. }
  107. def persist_search_payload(
  108. payload_json: str,
  109. *,
  110. run_id: str,
  111. keyword: str,
  112. query_reason: str,
  113. source_type: str,
  114. cursor: str,
  115. page_no: int,
  116. provider: str,
  117. source_value: str | None = None,
  118. parent_search_id: int | None = None,
  119. content_type: str = "视频",
  120. sort_type: str = "综合排序",
  121. publish_time: str = "不限",
  122. provider_state: dict[str, Any] | None = None,
  123. ) -> str:
  124. """新增搜索记录和本页全部候选,并把数据库 ID 拼回搜索结果。"""
  125. payload = _load_payload(payload_json)
  126. run_text = _clean_text(run_id, max_length=64)
  127. keyword_text = _clean_text(keyword, max_length=256)
  128. reason_text = _clean_text(query_reason)
  129. if not run_text or not keyword_text or not reason_text:
  130. payload["error"] = "run_id、keyword、query_reason 不能为空"
  131. payload["input_error"] = True
  132. return json.dumps(payload, ensure_ascii=False, default=str)
  133. if source_type not in _SOURCE_TYPES:
  134. payload["error"] = f"source_type 必须是: {sorted(_SOURCE_TYPES)}"
  135. payload["input_error"] = True
  136. return json.dumps(payload, ensure_ascii=False, default=str)
  137. raw_results = payload.get("search_results")
  138. results = raw_results if isinstance(raw_results, list) else []
  139. candidate_rows = [
  140. row
  141. for item in results
  142. if isinstance(item, dict)
  143. if (row := _candidate_from_search_result(dict(item), keyword_text)) is not None
  144. ]
  145. normalized_page_no = max(1, int(page_no))
  146. normalized_source_type = (
  147. "pagination" if normalized_page_no > 1 else source_type
  148. )
  149. normalized_parent_search_id = (
  150. None
  151. if normalized_source_type in _ROOT_SOURCE_TYPES
  152. else parent_search_id
  153. )
  154. merged_provider_state = dict(provider_state or {})
  155. for key in ("search_id", "backtrace"):
  156. value = payload.get(key)
  157. if value not in (None, ""):
  158. merged_provider_state[key] = value
  159. provider_search_id = payload.get("search_id")
  160. search_values: dict[str, Any] = {
  161. "run_id": run_text,
  162. "keyword": keyword_text,
  163. "query_reason": reason_text,
  164. "source_type": normalized_source_type,
  165. "source_value": _clean_text(source_value),
  166. "parent_search_id": normalized_parent_search_id,
  167. "provider": _clean_text(provider, max_length=32) or "internal_keyword",
  168. "provider_state_json": (
  169. json.dumps(merged_provider_state, ensure_ascii=False)
  170. if merged_provider_state
  171. else None
  172. ),
  173. "content_type": _clean_text(content_type, max_length=16) or "视频",
  174. "sort_type": _clean_text(sort_type, max_length=32) or "综合排序",
  175. "publish_time": _clean_text(publish_time, max_length=32) or "不限",
  176. "cursor": _clean_text(cursor, max_length=128) or "0",
  177. "page_no": normalized_page_no,
  178. "results_count": len(results),
  179. "new_candidate_count": 0,
  180. "has_more": int(bool(payload.get("has_more"))),
  181. "next_cursor": (
  182. _clean_text(payload.get("next_cursor"), max_length=128)
  183. ),
  184. "result_ids_json": None,
  185. "status": "failed" if payload.get("error") else "success",
  186. "error_message": _clean_text(payload.get("error")),
  187. }
  188. search_values["search_key"] = _search_key(search_values)
  189. try:
  190. saved = get_video_discovery_service().save_search_page(
  191. run_text,
  192. search_values,
  193. candidate_rows,
  194. )
  195. except RunNotFoundError as exc:
  196. payload["error"] = str(exc)
  197. payload["input_error"] = True
  198. return json.dumps(payload, ensure_ascii=False, default=str)
  199. except Exception as exc:
  200. logger.error("persist search payload failed: %s", exc, exc_info=True)
  201. payload["error"] = format_db_error(exc)
  202. return json.dumps(payload, ensure_ascii=False, default=str)
  203. payload.pop("search_results", None)
  204. payload.pop("user_videos", None)
  205. if provider_search_id not in (None, ""):
  206. payload["provider_search_id"] = provider_search_id
  207. payload.update(saved)
  208. payload["persisted"] = True
  209. return json.dumps(payload, ensure_ascii=False, default=str)