from __future__ import annotations import json import logging import os from datetime import datetime from typing import Any from zoneinfo import ZoneInfo import requests from supply_infra.config import get_infra_settings logger = logging.getLogger(__name__) AIGC_BASE_URL = "https://aigc-api.aiddit.com" CRAWLER_PLAN_CREATE_URL = f"{AIGC_BASE_URL}/aigc/crawler/plan/save" GET_PRODUCE_PLAN_DETAIL_BY_ID = f"{AIGC_BASE_URL}/aigc/produce/plan/detail" PRODUCE_PLAN_SAVE = f"{AIGC_BASE_URL}/aigc/produce/plan/save" DEFAULT_TIMEOUT = 60.0 SHANGHAI_TZ = ZoneInfo("Asia/Shanghai") MAX_VIDEOS_PER_CRAWLER_PLAN = 10 _INPUT_SOURCE_CHECK_KEYS = ("inputSourceModal", "inputSourceChannel", "contentType") class AigcClient: """AIGC 平台 HTTP 客户端(视频爬取计划 + 生成计划绑定)。""" def __init__(self, token: str | None = None) -> None: settings = get_infra_settings() self.token = ( token or os.getenv("AIGC_API_TOKEN") or settings.aigc_api_token or "" ).strip() def create_video_crawler_plan( self, aweme_ids: list[str], *, plan_name: str | None = None, ) -> dict[str, Any]: if not aweme_ids: raise ValueError("aweme_ids 不能为空") if len(aweme_ids) > MAX_VIDEOS_PER_CRAWLER_PLAN: raise ValueError( f"单次最多 {MAX_VIDEOS_PER_CRAWLER_PLAN} 个视频,当前 {len(aweme_ids)} 个" ) dt = datetime.now(SHANGHAI_TZ).strftime("%Y%m%d%H%M%S") crawler_plan_name = plan_name or f"【SupplyAgent】抖音视频直接抓取-{dt}-抖音" params = { "channel": 2, "contentModal": 4, "crawlerComment": 0, "crawlerMode": 5, "filterAccountMatchMode": 2, "filterContentMatchMode": 2, "frequencyType": 2, "inputModeValues": aweme_ids, "name": crawler_plan_name, "planType": 2, "searchModeValues": [], "srtExtractFlag": 1, "videoKeyFrameType": 1, "voiceExtractFlag": 1, } response_json = self._post(CRAWLER_PLAN_CREATE_URL, params) if response_json.get("code") != 0: message = response_json.get("msg", "创建爬取计划失败") return { "success": False, "error": message, "response": response_json, } crawler_plan_id = str(response_json.get("data", {}).get("id") or "").strip() return { "success": True, "crawler_plan_id": crawler_plan_id, "crawler_plan_name": crawler_plan_name, "aweme_ids": aweme_ids, } def bind_crawler_to_produce_plan( self, crawler_plan_id: str, produce_plan_id: str, *, crawler_plan_name: str, ) -> dict[str, Any]: if not crawler_plan_id or not produce_plan_id: raise ValueError("crawler_plan_id 与 produce_plan_id 均不能为空") input_source_info = { "contentType": 1, "inputSourceType": 2, "inputSourceValue": crawler_plan_id, "inputSourceLabel": f"原始帖子-视频-抖音-内容添加计划-{crawler_plan_name}", "inputSourceModal": 4, "inputSourceChannel": 2, } produce_plan_detail, error = self._get_produce_plan_detail(produce_plan_id) if error: return {"success": False, "produce_plan_id": produce_plan_id, "error": error} input_source_groups = produce_plan_detail.get("inputSourceGroups", []) if not input_source_groups: return { "success": False, "produce_plan_id": produce_plan_id, "error": "生成计划没有输入源组", } input_source_index = 0 for index, input_source_group in enumerate(input_source_groups): input_sources = input_source_group.get("inputSources", []) if not input_sources: continue first_input_source = input_sources[0] if all( input_source_info.get(key, 0) == first_input_source.get(key, -1) for key in _INPUT_SOURCE_CHECK_KEYS ): input_source_index = index break input_source_group = input_source_groups[input_source_index] input_source_group.setdefault("inputSources", []).append(input_source_info) response_json = self._post(PRODUCE_PLAN_SAVE, produce_plan_detail) if response_json.get("code") != 0 or not response_json.get("data", {}): return { "success": False, "produce_plan_id": produce_plan_id, "error": response_json.get("msg", "爬取计划绑定生成计划异常"), } return { "success": True, "produce_plan_id": produce_plan_id, "produce_plan_name": produce_plan_detail.get("name", ""), "msg": "成功", } def _get_produce_plan_detail(self, produce_plan_id: str) -> tuple[dict[str, Any], str | None]: response_json = self._post(GET_PRODUCE_PLAN_DETAIL_BY_ID, {"id": produce_plan_id}) if response_json.get("code") != 0 or not response_json.get("data", {}): return {}, response_json.get("msg", "获取生成计划详情异常") return response_json.get("data", {}), None def _post(self, url: str, params: Any) -> dict[str, Any]: if not self.token: raise RuntimeError("AIGC_API_TOKEN 未配置") request = { "baseInfo": {"token": self.token}, "params": params, } try: response = requests.post( url=url, json=request, headers={"Content-Type": "application/json"}, timeout=DEFAULT_TIMEOUT, ) response.raise_for_status() return response.json() except Exception as exc: logger.error( "invoke aigc platform error. url=%s params=%s error=%s", url, json.dumps(params, ensure_ascii=False), exc, ) return {}