classify.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344
  1. """帖子级「创作知识 / 非创作知识」分类 + 提取知识点。
  2. 判据是「内容作品创作知识」,覆盖图文、视频、脚本、游戏视频、历史视频等内容创作场景,
  3. 并显式排除三类越界(见 prompts):① 应试/学术写作
  4. ② 制作/工具操作(=制作知识)③ 学科知识/评论/作品本身。两套真实提示词:
  5. 图文(小红书/微信):prompts/classify_imgtext.txt(标题+正文+全部图喂多模态模型)
  6. 视频(抖音):prompts/classify_video.txt(OSS/CDN URL 直喂;本地大视频 fallback 时先压 480p 保音轨)
  7. 都输出 {is_empty, reason, knowledge}:is_empty 即创作闸;is_empty=false 时连带提炼出具体创作知识点。
  8. 正式链路只复用本模块的图文/视频判定函数;旧 SQLite 批量补判 CLI 已归档。
  9. 只读 prompts、走 OpenRouter / Ark / Qwen,不解构、不入 ingest。
  10. """
  11. from __future__ import annotations
  12. import base64
  13. import json
  14. import os
  15. import subprocess
  16. import tempfile
  17. import threading
  18. import time
  19. from pathlib import Path
  20. from urllib.parse import urlparse
  21. import httpx
  22. from core.config import Settings, load_env_file
  23. from core.prompts import load_prompt
  24. ROOT = Path(__file__).resolve().parent.parent
  25. PLATFORMS = ["xiaohongshu", "weixin", "douyin"] # 默认重判全部;可传平台名覆盖
  26. IMG_WORKERS = 8 # 图文并发
  27. VID_WORKERS = 3 # 视频并发(含 ffmpeg 压制,别太高)
  28. MAX_CARDS = 12 # 图文最多送几张图
  29. COMPRESS_OVER_MB = 12 # 视频超过此大小先压再喂
  30. COMPRESS_H = 480
  31. DEFAULT_PROVIDER_MIN_INTERVAL_SECONDS = 0.5
  32. DEFAULT_429_BACKOFF_SECONDS = [30.0, 90.0, 180.0]
  33. RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
  34. class _ProviderThrottle:
  35. def __init__(self) -> None:
  36. self._lock = threading.Lock()
  37. self._next_allowed = 0.0
  38. def wait(self, min_interval_seconds: float) -> None:
  39. with self._lock:
  40. now = time.monotonic()
  41. sleep_for = max(0.0, self._next_allowed - now)
  42. self._next_allowed = max(now, self._next_allowed) + min_interval_seconds
  43. if sleep_for > 0:
  44. time.sleep(sleep_for)
  45. def penalize(self, seconds: float) -> None:
  46. if seconds <= 0:
  47. return
  48. with self._lock:
  49. self._next_allowed = max(self._next_allowed, time.monotonic() + seconds)
  50. _PROVIDER_THROTTLES: dict[str, _ProviderThrottle] = {}
  51. _PROVIDER_THROTTLES_LOCK = threading.Lock()
  52. def _data_root(settings: Settings) -> Path:
  53. p = Path(settings.data_dir or "data")
  54. return p if p.is_absolute() else ROOT / p
  55. def _public_path_to_file(public_path: str, settings: Settings) -> Path:
  56. if public_path.startswith("/data/"):
  57. return _data_root(settings) / public_path.removeprefix("/data/")
  58. return ROOT / public_path.lstrip("/")
  59. def _data_url(public_path: str, settings: Settings):
  60. fs = _public_path_to_file(public_path, settings)
  61. if not fs.exists():
  62. return None
  63. return "data:image/jpeg;base64," + base64.b64encode(fs.read_bytes()).decode()
  64. def _is_http_url(value: str) -> bool:
  65. try:
  66. parsed = urlparse(value)
  67. except Exception:
  68. return False
  69. return parsed.scheme in ("http", "https") and bool(parsed.netloc)
  70. def _compress(mp4: Path) -> Path:
  71. """大视频压到 480p(保留口播音轨)→ 临时 mp4;失败回原文件。"""
  72. out = Path(tempfile.gettempdir()) / f"ck_{mp4.parent.name}_480.mp4"
  73. try:
  74. import imageio_ffmpeg
  75. ff = imageio_ffmpeg.get_ffmpeg_exe()
  76. subprocess.run([ff, "-y", "-i", str(mp4), "-vf", f"scale=-2:{COMPRESS_H}",
  77. "-c:v", "libx264", "-crf", "30", "-preset", "veryfast",
  78. "-c:a", "aac", "-b:a", "64k", str(out)],
  79. capture_output=True, timeout=300)
  80. if out.exists() and out.stat().st_size > 0:
  81. return out
  82. except Exception:
  83. pass
  84. return mp4
  85. def _parse_judge_content(content: str) -> tuple:
  86. d = json.loads(content)
  87. if bool(d.get("is_empty")):
  88. return 0, str(d.get("reason", ""))[:60], "", ""
  89. return 1, str(d.get("reason", ""))[:60], str(d.get("knowledge", "") or ""), ""
  90. def _env_first(env: dict, *keys: str) -> str:
  91. for key in keys:
  92. value = os.getenv(key) or env.get(key)
  93. if value:
  94. return value
  95. return ""
  96. def _is_qwen_model(model: str) -> bool:
  97. lower = model.lower()
  98. return lower.startswith("qwen") or lower.startswith("qwq")
  99. def _provider_key(name: str) -> str:
  100. return name.split(":", 1)[0].lower()
  101. def _provider_throttle(provider: str) -> _ProviderThrottle:
  102. with _PROVIDER_THROTTLES_LOCK:
  103. throttle = _PROVIDER_THROTTLES.get(provider)
  104. if throttle is None:
  105. throttle = _ProviderThrottle()
  106. _PROVIDER_THROTTLES[provider] = throttle
  107. return throttle
  108. def _env_float(env: dict, default: float, *keys: str) -> float:
  109. value = _env_first(env, *keys)
  110. if not value:
  111. return default
  112. try:
  113. return max(0.0, float(value))
  114. except ValueError:
  115. return default
  116. def _provider_min_interval(provider: str, env: dict) -> float:
  117. prefix = provider.upper()
  118. return _env_float(
  119. env,
  120. DEFAULT_PROVIDER_MIN_INTERVAL_SECONDS,
  121. f"CLASSIFY_{prefix}_MIN_INTERVAL_SECONDS",
  122. "CLASSIFY_PROVIDER_MIN_INTERVAL_SECONDS",
  123. )
  124. def _parse_backoff_list(value: str) -> list[float]:
  125. out = []
  126. for part in value.split(","):
  127. try:
  128. seconds = float(part.strip())
  129. except ValueError:
  130. continue
  131. if seconds > 0:
  132. out.append(seconds)
  133. return out or DEFAULT_429_BACKOFF_SECONDS
  134. def _retry_after_seconds(resp: httpx.Response) -> float | None:
  135. value = resp.headers.get("retry-after")
  136. if not value:
  137. return None
  138. try:
  139. return max(0.0, float(value))
  140. except ValueError:
  141. return None
  142. def _provider_429_backoff(provider: str, env: dict, attempt: int, resp: httpx.Response) -> float:
  143. retry_after = _retry_after_seconds(resp)
  144. if retry_after is not None:
  145. return retry_after
  146. prefix = provider.upper()
  147. values = _parse_backoff_list(
  148. _env_first(
  149. env,
  150. f"CLASSIFY_{prefix}_429_BACKOFF_SECONDS",
  151. "CLASSIFY_429_BACKOFF_SECONDS",
  152. ) or ",".join(str(v) for v in DEFAULT_429_BACKOFF_SECONDS)
  153. )
  154. return values[min(attempt, len(values) - 1)]
  155. def _providers(settings: Settings, messages: list) -> list[tuple[str, str, dict, dict]]:
  156. body = {"model": settings.video_model, "messages": messages,
  157. "response_format": {"type": "json_object"}}
  158. env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
  159. out: list[tuple[str, str, dict, dict]] = []
  160. provider = (_env_first(env, "CLASSIFY_PROVIDER") or "auto").lower()
  161. explicit_model = _env_first(env, "CLASSIFY_MODEL")
  162. bailian_key = _env_first(
  163. env,
  164. "ALIYUN_BAILIAN_API_KEY",
  165. )
  166. if bailian_key and provider in ("auto", "qwen", "bailian"):
  167. bailian_url = _env_first(
  168. env,
  169. "ALIYUN_BAILIAN_BASE_URL",
  170. ) or "https://dashscope.aliyuncs.com/compatible-mode/v1"
  171. models = []
  172. for model in [
  173. explicit_model,
  174. _env_first(env, "ALIYUN_BAILIAN_MODEL"),
  175. _env_first(env, "OPENROUTER_MODEL"),
  176. "qwen3.7-plus",
  177. "qwen-vl-plus",
  178. ]:
  179. if model and _is_qwen_model(model) and model not in models:
  180. models.append(model)
  181. for model in models:
  182. out.append((
  183. f"qwen:{model}",
  184. bailian_url.rstrip("/") + "/chat/completions",
  185. {"Authorization": f"Bearer {bailian_key}", "Content-Type": "application/json"},
  186. {"model": model, "messages": messages, "response_format": {"type": "json_object"}},
  187. ))
  188. if provider in ("qwen", "bailian"):
  189. return out
  190. if settings.openrouter_api_key and provider in ("auto", "openrouter"):
  191. out.append((
  192. "openrouter",
  193. settings.openrouter_base_url.rstrip("/") + "/chat/completions",
  194. {"Authorization": f"Bearer {settings.openrouter_api_key}", "Content-Type": "application/json"},
  195. body,
  196. ))
  197. ark_key = _env_first(env, "ARK_API_KEY")
  198. if ark_key and provider in ("auto", "ark"):
  199. ark_url = _env_first(env, "ARK_CHAT_URL") or "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
  200. models = []
  201. explicit = explicit_model or _env_first(env, "ARK_CHAT_MODEL")
  202. if explicit:
  203. models.append(explicit)
  204. # Seed 2 Mini 接入点/1.6 Vision 更适合图文+视频理解;flash 作为轻量兜底。
  205. models.extend(["ep-20260506151915-jqvw7", "doubao-seed-1-6-vision-250815", "doubao-seed-1-6-flash-250615"])
  206. seen = set()
  207. for model in models:
  208. if model in seen:
  209. continue
  210. seen.add(model)
  211. out.append((
  212. f"ark:{model}",
  213. ark_url,
  214. {"Authorization": f"Bearer {ark_key}", "Content-Type": "application/json"},
  215. {"model": model, "messages": messages, "response_format": {"type": "json_object"}},
  216. ))
  217. return out
  218. def _judge(messages: list, settings: Settings, timeout: float) -> tuple:
  219. """调多模态模型(Qwen / OpenRouter / Ark),解析 {is_empty, reason, knowledge},带重试。
  220. 返回 (is_creation 1/0/None, reason, knowledge, points)。"""
  221. last = ""
  222. env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
  223. for name, api, headers, payload in _providers(settings, messages):
  224. provider = _provider_key(name)
  225. throttle = _provider_throttle(provider)
  226. min_interval = _provider_min_interval(provider, env)
  227. for attempt in range(3):
  228. try:
  229. throttle.wait(min_interval)
  230. resp = httpx.post(api, headers=headers, json=payload, timeout=timeout)
  231. if resp.status_code == 200:
  232. return _parse_judge_content(resp.json()["choices"][0]["message"]["content"])
  233. last = f"{name} http {resp.status_code}"
  234. if resp.status_code == 429:
  235. throttle.penalize(_provider_429_backoff(provider, env, attempt, resp))
  236. if resp.status_code in (401, 403, 404):
  237. break
  238. except Exception as exc:
  239. last = f"{name} {str(exc)[:50]}"
  240. if "http " not in last or any(f"http {code}" in last for code in RETRYABLE_STATUS_CODES):
  241. time.sleep(2 * (attempt + 1))
  242. else:
  243. break
  244. if not last:
  245. last = "missing Qwen/OpenRouter/Ark credentials"
  246. return None, f"判定失败: {last}", "", ""
  247. def classify_imgtext(p: dict, settings: Settings) -> tuple:
  248. """图文:标题+正文+全部图,用收紧的 classify_imgtext.txt 判 is_empty 并提取知识点。"""
  249. user = [{"type": "text", "text": f"平台:{p.get('platform')}\n标题:{p.get('title', '')}\n"
  250. f"正文:{(p.get('body_text') or '')[:1500]}\n(下附帖子图片,请一并看完)"}]
  251. for im in (p.get("images") or [])[:MAX_CARDS]:
  252. if _is_http_url(im):
  253. user.append({"type": "image_url", "image_url": {"url": im}})
  254. continue
  255. u = _data_url(im, settings)
  256. if u:
  257. user.append({"type": "image_url", "image_url": {"url": u}})
  258. messages = [{"role": "system", "content": load_prompt("classify_imgtext")},
  259. {"role": "user", "content": user}]
  260. return _judge(messages, settings, timeout=120)
  261. def classify_video(p: dict, settings: Settings) -> tuple:
  262. """视频:看完整段视频,用 classify_video.txt 判 is_empty 并提取知识点。
  263. 新采集链路里抖音视频会先转存 OSS,传入 HTTP(S) CDN URL;老链路传 /data 本地 mp4。
  264. """
  265. rel = p.get("video") or ""
  266. if _is_http_url(rel):
  267. media = rel
  268. else:
  269. mp4 = _public_path_to_file(rel, settings)
  270. if not rel or not mp4.exists():
  271. return None, "无视频", "", ""
  272. use = _compress(mp4) if mp4.stat().st_size > COMPRESS_OVER_MB * 1048576 else mp4
  273. try:
  274. media = "data:video/mp4;base64," + base64.b64encode(use.read_bytes()).decode()
  275. finally:
  276. if use != mp4:
  277. try:
  278. use.unlink()
  279. except Exception:
  280. pass
  281. messages = [{"role": "system", "content": load_prompt("classify_video")},
  282. {"role": "user", "content": [{"type": "text", "text": "判断这条视频是不是创作知识。"},
  283. {"type": "video_url", "video_url": {"url": media}}]}]
  284. return _judge(messages, settings, timeout=300)
  285. if __name__ == "__main__":
  286. raise SystemExit(
  287. "acquisition.classify now exposes classifier functions only; "
  288. "use acquisition/classification/coarse.py or the formal acquisition runner."
  289. )