search_persistence.py 8.4 KB

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