client.py 7.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201
  1. from __future__ import annotations
  2. import json
  3. import logging
  4. import os
  5. from datetime import datetime
  6. from typing import Any
  7. from zoneinfo import ZoneInfo
  8. import requests
  9. from supply_infra.config import get_infra_settings
  10. logger = logging.getLogger(__name__)
  11. AIGC_BASE_URL = "https://aigc-api.aiddit.com"
  12. CRAWLER_PLAN_CREATE_URL = f"{AIGC_BASE_URL}/aigc/crawler/plan/save"
  13. GET_PRODUCE_PLAN_DETAIL_BY_ID = f"{AIGC_BASE_URL}/aigc/produce/plan/detail"
  14. PRODUCE_PLAN_SAVE = f"{AIGC_BASE_URL}/aigc/produce/plan/save"
  15. DEFAULT_TIMEOUT = 60.0
  16. SHANGHAI_TZ = ZoneInfo("Asia/Shanghai")
  17. MAX_VIDEOS_PER_CRAWLER_PLAN = 10
  18. _INPUT_SOURCE_CHECK_KEYS = ("inputSourceModal", "inputSourceChannel", "contentType")
  19. class AigcClient:
  20. """AIGC 平台 HTTP 客户端(视频爬取计划 + 生成计划绑定)。"""
  21. def __init__(self, token: str | None = None) -> None:
  22. settings = get_infra_settings()
  23. resolved_token = (
  24. token
  25. if token is not None
  26. else (
  27. os.getenv("AIGC_API_TOKEN")
  28. or settings.aigc_api_token
  29. or ""
  30. )
  31. )
  32. self.token = resolved_token.strip()
  33. def create_video_crawler_plan(
  34. self,
  35. aweme_ids: list[str],
  36. *,
  37. plan_name: str | None = None,
  38. ) -> dict[str, Any]:
  39. if not aweme_ids:
  40. raise ValueError("aweme_ids 不能为空")
  41. if len(aweme_ids) > MAX_VIDEOS_PER_CRAWLER_PLAN:
  42. raise ValueError(
  43. f"单次最多 {MAX_VIDEOS_PER_CRAWLER_PLAN} 个视频,当前 {len(aweme_ids)} 个"
  44. )
  45. dt = datetime.now(SHANGHAI_TZ).strftime("%Y%m%d%H%M%S")
  46. crawler_plan_name = plan_name or f"【SupplyAgent】抖音视频直接抓取-{dt}-抖音"
  47. params = {
  48. "channel": 2,
  49. "contentModal": 4,
  50. "crawlerComment": 0,
  51. "crawlerMode": 5,
  52. "filterAccountMatchMode": 2,
  53. "filterContentMatchMode": 2,
  54. "frequencyType": 2,
  55. "inputModeValues": aweme_ids,
  56. "name": crawler_plan_name,
  57. "planType": 2,
  58. "searchModeValues": [],
  59. "srtExtractFlag": 1,
  60. "videoKeyFrameType": 1,
  61. "voiceExtractFlag": 1,
  62. }
  63. response_json = self._post(CRAWLER_PLAN_CREATE_URL, params)
  64. if response_json.get("code") != 0:
  65. message = response_json.get("msg", "创建爬取计划失败")
  66. return {
  67. "success": False,
  68. "error": message,
  69. "response": response_json,
  70. }
  71. crawler_plan_id = str(response_json.get("data", {}).get("id") or "").strip()
  72. return {
  73. "success": True,
  74. "crawler_plan_id": crawler_plan_id,
  75. "crawler_plan_name": crawler_plan_name,
  76. "aweme_ids": aweme_ids,
  77. }
  78. def bind_crawler_to_produce_plan(
  79. self,
  80. crawler_plan_id: str,
  81. produce_plan_id: str,
  82. *,
  83. crawler_plan_name: str,
  84. ) -> dict[str, Any]:
  85. if not crawler_plan_id or not produce_plan_id:
  86. raise ValueError("crawler_plan_id 与 produce_plan_id 均不能为空")
  87. input_source_info = {
  88. "contentType": 1,
  89. "inputSourceType": 2,
  90. "inputSourceValue": crawler_plan_id,
  91. "inputSourceLabel": f"原始帖子-视频-抖音-内容添加计划-{crawler_plan_name}",
  92. "inputSourceModal": 4,
  93. "inputSourceChannel": 2,
  94. }
  95. produce_plan_detail, error = self._get_produce_plan_detail(produce_plan_id)
  96. if error:
  97. return {"success": False, "produce_plan_id": produce_plan_id, "error": error}
  98. input_source_groups = produce_plan_detail.get("inputSourceGroups", [])
  99. if not input_source_groups:
  100. return {
  101. "success": False,
  102. "produce_plan_id": produce_plan_id,
  103. "error": "生成计划没有输入源组",
  104. }
  105. for input_source_group in input_source_groups:
  106. for existing_source in input_source_group.get("inputSources", []):
  107. same_crawler = (
  108. str(existing_source.get("inputSourceValue") or "")
  109. == crawler_plan_id
  110. )
  111. same_shape = all(
  112. input_source_info.get(key, 0) == existing_source.get(key, -1)
  113. for key in _INPUT_SOURCE_CHECK_KEYS
  114. )
  115. if same_crawler and same_shape:
  116. return {
  117. "success": True,
  118. "already_bound": True,
  119. "produce_plan_id": produce_plan_id,
  120. "produce_plan_name": produce_plan_detail.get("name", ""),
  121. "msg": "爬取计划已绑定,无需重复写入",
  122. }
  123. input_source_index = 0
  124. for index, input_source_group in enumerate(input_source_groups):
  125. input_sources = input_source_group.get("inputSources", [])
  126. if not input_sources:
  127. continue
  128. first_input_source = input_sources[0]
  129. if all(
  130. input_source_info.get(key, 0) == first_input_source.get(key, -1)
  131. for key in _INPUT_SOURCE_CHECK_KEYS
  132. ):
  133. input_source_index = index
  134. break
  135. input_source_group = input_source_groups[input_source_index]
  136. input_source_group.setdefault("inputSources", []).append(input_source_info)
  137. response_json = self._post(PRODUCE_PLAN_SAVE, produce_plan_detail)
  138. if response_json.get("code") != 0 or not response_json.get("data", {}):
  139. return {
  140. "success": False,
  141. "produce_plan_id": produce_plan_id,
  142. "error": response_json.get("msg", "爬取计划绑定生成计划异常"),
  143. }
  144. return {
  145. "success": True,
  146. "produce_plan_id": produce_plan_id,
  147. "produce_plan_name": produce_plan_detail.get("name", ""),
  148. "msg": "成功",
  149. }
  150. def _get_produce_plan_detail(self, produce_plan_id: str) -> tuple[dict[str, Any], str | None]:
  151. response_json = self._post(GET_PRODUCE_PLAN_DETAIL_BY_ID, {"id": produce_plan_id})
  152. if response_json.get("code") != 0 or not response_json.get("data", {}):
  153. return {}, response_json.get("msg", "获取生成计划详情异常")
  154. return response_json.get("data", {}), None
  155. def _post(self, url: str, params: Any) -> dict[str, Any]:
  156. if not self.token:
  157. raise RuntimeError("AIGC_API_TOKEN 未配置")
  158. request = {
  159. "baseInfo": {"token": self.token},
  160. "params": params,
  161. }
  162. try:
  163. response = requests.post(
  164. url=url,
  165. json=request,
  166. headers={"Content-Type": "application/json"},
  167. timeout=DEFAULT_TIMEOUT,
  168. )
  169. response.raise_for_status()
  170. return response.json()
  171. except Exception as exc:
  172. logger.error(
  173. "invoke aigc platform error. url=%s params=%s error=%s",
  174. url,
  175. json.dumps(params, ensure_ascii=False),
  176. exc,
  177. )
  178. return {}