classify.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460
  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. from core.text_limits import (
  25. CLASSIFY_BODY_MAX_CHARS,
  26. ERROR_MESSAGE_MAX_CHARS,
  27. REASON_MAX_CHARS,
  28. clip_text,
  29. )
  30. from pipeline.tracing import TraceContext, TraceWriter, hash_prompt, redact_headers, timed_ms
  31. ROOT = Path(__file__).resolve().parent.parent
  32. PLATFORMS = ["xiaohongshu", "weixin", "douyin"] # 默认重判全部;可传平台名覆盖
  33. IMG_WORKERS = 8 # 图文并发
  34. VID_WORKERS = 3 # 视频并发(含 ffmpeg 压制,别太高)
  35. MAX_CARDS = 12 # 图文最多送几张图
  36. COMPRESS_OVER_MB = 12 # 视频超过此大小先压再喂
  37. COMPRESS_H = 480
  38. DEFAULT_PROVIDER_MIN_INTERVAL_SECONDS = 0.5
  39. DEFAULT_429_BACKOFF_SECONDS = [30.0, 90.0, 180.0]
  40. RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
  41. class _ProviderThrottle:
  42. def __init__(self) -> None:
  43. self._lock = threading.Lock()
  44. self._next_allowed = 0.0
  45. def wait(self, min_interval_seconds: float) -> None:
  46. with self._lock:
  47. now = time.monotonic()
  48. sleep_for = max(0.0, self._next_allowed - now)
  49. self._next_allowed = max(now, self._next_allowed) + min_interval_seconds
  50. if sleep_for > 0:
  51. time.sleep(sleep_for)
  52. def penalize(self, seconds: float) -> None:
  53. if seconds <= 0:
  54. return
  55. with self._lock:
  56. self._next_allowed = max(self._next_allowed, time.monotonic() + seconds)
  57. _PROVIDER_THROTTLES: dict[str, _ProviderThrottle] = {}
  58. _PROVIDER_THROTTLES_LOCK = threading.Lock()
  59. def _data_root(settings: Settings) -> Path:
  60. p = Path(settings.data_dir or "data")
  61. return p if p.is_absolute() else ROOT / p
  62. def _public_path_to_file(public_path: str, settings: Settings) -> Path:
  63. if public_path.startswith("/data/"):
  64. return _data_root(settings) / public_path.removeprefix("/data/")
  65. return ROOT / public_path.lstrip("/")
  66. def _data_url(public_path: str, settings: Settings):
  67. fs = _public_path_to_file(public_path, settings)
  68. if not fs.exists():
  69. return None
  70. return "data:image/jpeg;base64," + base64.b64encode(fs.read_bytes()).decode()
  71. def _is_http_url(value: str) -> bool:
  72. try:
  73. parsed = urlparse(value)
  74. except Exception:
  75. return False
  76. return parsed.scheme in ("http", "https") and bool(parsed.netloc)
  77. def _compress(mp4: Path) -> Path:
  78. """大视频压到 480p(保留口播音轨)→ 临时 mp4;失败回原文件。"""
  79. out = Path(tempfile.gettempdir()) / f"ck_{mp4.parent.name}_480.mp4"
  80. try:
  81. import imageio_ffmpeg
  82. ff = imageio_ffmpeg.get_ffmpeg_exe()
  83. subprocess.run([ff, "-y", "-i", str(mp4), "-vf", f"scale=-2:{COMPRESS_H}",
  84. "-c:v", "libx264", "-crf", "30", "-preset", "veryfast",
  85. "-c:a", "aac", "-b:a", "64k", str(out)],
  86. capture_output=True, timeout=300)
  87. if out.exists() and out.stat().st_size > 0:
  88. return out
  89. except Exception:
  90. pass
  91. return mp4
  92. def _parse_judge_content(content: str) -> tuple:
  93. d = json.loads(content)
  94. if bool(d.get("is_empty")):
  95. return 0, clip_text(d.get("reason", ""), REASON_MAX_CHARS), "", ""
  96. return 1, clip_text(d.get("reason", ""), REASON_MAX_CHARS), str(d.get("knowledge", "") or ""), ""
  97. def _env_first(env: dict, *keys: str) -> str:
  98. for key in keys:
  99. value = os.getenv(key) or env.get(key)
  100. if value:
  101. return value
  102. return ""
  103. def _is_qwen_model(model: str) -> bool:
  104. lower = model.lower()
  105. return lower.startswith("qwen") or lower.startswith("qwq")
  106. def _provider_key(name: str) -> str:
  107. return name.split(":", 1)[0].lower()
  108. def _provider_throttle(provider: str) -> _ProviderThrottle:
  109. with _PROVIDER_THROTTLES_LOCK:
  110. throttle = _PROVIDER_THROTTLES.get(provider)
  111. if throttle is None:
  112. throttle = _ProviderThrottle()
  113. _PROVIDER_THROTTLES[provider] = throttle
  114. return throttle
  115. def _env_float(env: dict, default: float, *keys: str) -> float:
  116. value = _env_first(env, *keys)
  117. if not value:
  118. return default
  119. try:
  120. return max(0.0, float(value))
  121. except ValueError:
  122. return default
  123. def _provider_min_interval(provider: str, env: dict) -> float:
  124. prefix = provider.upper()
  125. return _env_float(
  126. env,
  127. DEFAULT_PROVIDER_MIN_INTERVAL_SECONDS,
  128. f"CLASSIFY_{prefix}_MIN_INTERVAL_SECONDS",
  129. "CLASSIFY_PROVIDER_MIN_INTERVAL_SECONDS",
  130. )
  131. def _parse_backoff_list(value: str) -> list[float]:
  132. out = []
  133. for part in value.split(","):
  134. try:
  135. seconds = float(part.strip())
  136. except ValueError:
  137. continue
  138. if seconds > 0:
  139. out.append(seconds)
  140. return out or DEFAULT_429_BACKOFF_SECONDS
  141. def _retry_after_seconds(resp: httpx.Response) -> float | None:
  142. value = resp.headers.get("retry-after")
  143. if not value:
  144. return None
  145. try:
  146. return max(0.0, float(value))
  147. except ValueError:
  148. return None
  149. def _provider_429_backoff(provider: str, env: dict, attempt: int, resp: httpx.Response) -> float:
  150. retry_after = _retry_after_seconds(resp)
  151. if retry_after is not None:
  152. return retry_after
  153. prefix = provider.upper()
  154. values = _parse_backoff_list(
  155. _env_first(
  156. env,
  157. f"CLASSIFY_{prefix}_429_BACKOFF_SECONDS",
  158. "CLASSIFY_429_BACKOFF_SECONDS",
  159. ) or ",".join(str(v) for v in DEFAULT_429_BACKOFF_SECONDS)
  160. )
  161. return values[min(attempt, len(values) - 1)]
  162. def _providers(settings: Settings, messages: list) -> list[tuple[str, str, dict, dict]]:
  163. body = {"model": settings.video_model, "messages": messages,
  164. "response_format": {"type": "json_object"}}
  165. env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
  166. out: list[tuple[str, str, dict, dict]] = []
  167. provider = (_env_first(env, "CLASSIFY_PROVIDER") or "auto").lower()
  168. explicit_model = _env_first(env, "CLASSIFY_MODEL")
  169. bailian_key = _env_first(
  170. env,
  171. "ALIYUN_BAILIAN_API_KEY",
  172. )
  173. if bailian_key and provider in ("auto", "qwen", "bailian"):
  174. bailian_url = _env_first(
  175. env,
  176. "ALIYUN_BAILIAN_BASE_URL",
  177. ) or "https://dashscope.aliyuncs.com/compatible-mode/v1"
  178. models = []
  179. for model in [
  180. explicit_model,
  181. _env_first(env, "ALIYUN_BAILIAN_MODEL"),
  182. _env_first(env, "OPENROUTER_MODEL"),
  183. "qwen3.7-plus",
  184. "qwen-vl-plus",
  185. ]:
  186. if model and _is_qwen_model(model) and model not in models:
  187. models.append(model)
  188. for model in models:
  189. out.append((
  190. f"qwen:{model}",
  191. bailian_url.rstrip("/") + "/chat/completions",
  192. {"Authorization": f"Bearer {bailian_key}", "Content-Type": "application/json"},
  193. {"model": model, "messages": messages, "response_format": {"type": "json_object"}},
  194. ))
  195. if provider in ("qwen", "bailian"):
  196. return out
  197. if settings.openrouter_api_key and provider in ("auto", "openrouter"):
  198. out.append((
  199. "openrouter",
  200. settings.openrouter_base_url.rstrip("/") + "/chat/completions",
  201. {"Authorization": f"Bearer {settings.openrouter_api_key}", "Content-Type": "application/json"},
  202. body,
  203. ))
  204. ark_key = _env_first(env, "ARK_API_KEY")
  205. if ark_key and provider in ("auto", "ark"):
  206. ark_url = _env_first(env, "ARK_CHAT_URL") or "https://ark.cn-beijing.volces.com/api/v3/chat/completions"
  207. models = []
  208. explicit = explicit_model or _env_first(env, "ARK_CHAT_MODEL")
  209. if explicit:
  210. models.append(explicit)
  211. # Seed 2 Mini 接入点/1.6 Vision 更适合图文+视频理解;flash 作为轻量兜底。
  212. models.extend(["ep-20260506151915-jqvw7", "doubao-seed-1-6-vision-250815", "doubao-seed-1-6-flash-250615"])
  213. seen = set()
  214. for model in models:
  215. if model in seen:
  216. continue
  217. seen.add(model)
  218. out.append((
  219. f"ark:{model}",
  220. ark_url,
  221. {"Authorization": f"Bearer {ark_key}", "Content-Type": "application/json"},
  222. {"model": model, "messages": messages, "response_format": {"type": "json_object"}},
  223. ))
  224. return out
  225. def _judge(
  226. messages: list,
  227. settings: Settings,
  228. timeout: float,
  229. *,
  230. trace_writer: TraceWriter | None = None,
  231. trace_context: TraceContext | None = None,
  232. trace_stage: str = "classify",
  233. trace_substage: str = "coarse",
  234. prompt_name: str = "classify",
  235. ) -> tuple:
  236. """调多模态模型(Qwen / OpenRouter / Ark),解析 {is_empty, reason, knowledge},带重试。
  237. 返回 (is_creation 1/0/None, reason, knowledge, points)。"""
  238. last = ""
  239. env = load_env_file(os.getenv("CK_ENV_FILE", ".env"))
  240. for name, api, headers, payload in _providers(settings, messages):
  241. provider = _provider_key(name)
  242. throttle = _provider_throttle(provider)
  243. min_interval = _provider_min_interval(provider, env)
  244. for attempt in range(3):
  245. started = time.perf_counter()
  246. try:
  247. throttle.wait(min_interval)
  248. resp = httpx.post(api, headers=headers, json=payload, timeout=timeout)
  249. if resp.status_code == 200:
  250. response_json = resp.json()
  251. content = response_json["choices"][0]["message"]["content"]
  252. parsed = _parse_judge_content(content)
  253. if trace_writer is not None:
  254. trace_writer.llm_call(
  255. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  256. stage=trace_stage,
  257. substage=trace_substage,
  258. provider=name,
  259. model_name=payload.get("model"),
  260. endpoint=api,
  261. prompt_name=prompt_name,
  262. prompt_hash=hash_prompt(str(messages[0].get("content") if messages else "")),
  263. request_payload={"headers": redact_headers(headers), **payload},
  264. response_payload=response_json,
  265. parsed_payload={
  266. "is_creation": parsed[0],
  267. "reason": parsed[1],
  268. "knowledge": parsed[2],
  269. "points": parsed[3],
  270. },
  271. status="done",
  272. latency_ms=timed_ms(started),
  273. attempt_index=attempt + 1,
  274. )
  275. return parsed
  276. last = f"{name} http {resp.status_code}"
  277. if trace_writer is not None:
  278. trace_writer.llm_call(
  279. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  280. stage=trace_stage,
  281. substage=trace_substage,
  282. provider=name,
  283. model_name=payload.get("model"),
  284. endpoint=api,
  285. prompt_name=prompt_name,
  286. prompt_hash=hash_prompt(str(messages[0].get("content") if messages else "")),
  287. request_payload={"headers": redact_headers(headers), **payload},
  288. response_payload={"http_status": resp.status_code, "text": resp.text},
  289. status="failed",
  290. error_message=last,
  291. latency_ms=timed_ms(started),
  292. attempt_index=attempt + 1,
  293. )
  294. if resp.status_code == 429:
  295. throttle.penalize(_provider_429_backoff(provider, env, attempt, resp))
  296. if resp.status_code in (401, 403, 404):
  297. break
  298. except Exception as exc:
  299. last = f"{name} {clip_text(exc, ERROR_MESSAGE_MAX_CHARS)}"
  300. if trace_writer is not None:
  301. trace_writer.llm_call(
  302. context=trace_context or TraceContext(stage=trace_stage, substage=trace_substage),
  303. stage=trace_stage,
  304. substage=trace_substage,
  305. provider=name,
  306. model_name=payload.get("model"),
  307. endpoint=api,
  308. prompt_name=prompt_name,
  309. prompt_hash=hash_prompt(str(messages[0].get("content") if messages else "")),
  310. request_payload={"headers": redact_headers(headers), **payload},
  311. status="failed",
  312. error_message=last,
  313. latency_ms=timed_ms(started),
  314. attempt_index=attempt + 1,
  315. )
  316. if "http " not in last or any(f"http {code}" in last for code in RETRYABLE_STATUS_CODES):
  317. time.sleep(2 * (attempt + 1))
  318. else:
  319. break
  320. if not last:
  321. last = "missing Qwen/OpenRouter/Ark credentials"
  322. return None, f"判定失败: {last}", "", ""
  323. def classify_imgtext(
  324. p: dict,
  325. settings: Settings,
  326. *,
  327. trace_writer: TraceWriter | None = None,
  328. trace_context: TraceContext | None = None,
  329. ) -> tuple:
  330. """图文:标题+正文+全部图,用收紧的 classify_imgtext.txt 判 is_empty 并提取知识点。"""
  331. user = [{"type": "text", "text": f"平台:{p.get('platform')}\n标题:{p.get('title', '')}\n"
  332. f"正文:{clip_text(p.get('body_text') or '', CLASSIFY_BODY_MAX_CHARS)}\n(下附帖子图片,请一并看完)"}]
  333. for im in (p.get("images") or [])[:MAX_CARDS]:
  334. if _is_http_url(im):
  335. user.append({"type": "image_url", "image_url": {"url": im}})
  336. continue
  337. u = _data_url(im, settings)
  338. if u:
  339. user.append({"type": "image_url", "image_url": {"url": u}})
  340. messages = [{"role": "system", "content": load_prompt("classify_imgtext")},
  341. {"role": "user", "content": user}]
  342. if trace_writer is None and trace_context is None:
  343. return _judge(messages, settings, timeout=120)
  344. return _judge(
  345. messages,
  346. settings,
  347. timeout=120,
  348. trace_writer=trace_writer,
  349. trace_context=trace_context,
  350. trace_stage="classify",
  351. trace_substage="coarse_imgtext",
  352. prompt_name="classify_imgtext",
  353. )
  354. def classify_video(
  355. p: dict,
  356. settings: Settings,
  357. *,
  358. trace_writer: TraceWriter | None = None,
  359. trace_context: TraceContext | None = None,
  360. ) -> tuple:
  361. """视频:看完整段视频,用 classify_video.txt 判 is_empty 并提取知识点。
  362. 新采集链路里抖音视频会先转存 OSS,传入 HTTP(S) CDN URL;老链路传 /data 本地 mp4。
  363. """
  364. rel = p.get("video") or ""
  365. if _is_http_url(rel):
  366. media = rel
  367. else:
  368. mp4 = _public_path_to_file(rel, settings)
  369. if not rel or not mp4.exists():
  370. return None, "无视频", "", ""
  371. use = _compress(mp4) if mp4.stat().st_size > COMPRESS_OVER_MB * 1048576 else mp4
  372. try:
  373. media = "data:video/mp4;base64," + base64.b64encode(use.read_bytes()).decode()
  374. finally:
  375. if use != mp4:
  376. try:
  377. use.unlink()
  378. except Exception:
  379. pass
  380. user_text = (
  381. "判断这条视频是不是创作知识。\n"
  382. f"平台:{p.get('platform', '')}\n"
  383. f"标题:{p.get('title', '')}\n"
  384. f"正文/文案:{clip_text(p.get('body_text') or '', CLASSIFY_BODY_MAX_CHARS)}"
  385. )
  386. messages = [{"role": "system", "content": load_prompt("classify_video")},
  387. {"role": "user", "content": [{"type": "text", "text": user_text},
  388. {"type": "video_url", "video_url": {"url": media}}]}]
  389. if trace_writer is None and trace_context is None:
  390. return _judge(messages, settings, timeout=300)
  391. return _judge(
  392. messages,
  393. settings,
  394. timeout=300,
  395. trace_writer=trace_writer,
  396. trace_context=trace_context,
  397. trace_stage="classify",
  398. trace_substage="coarse_video",
  399. prompt_name="classify_video",
  400. )
  401. if __name__ == "__main__":
  402. raise SystemExit(
  403. "acquisition.classify now exposes classifier functions only; "
  404. "use acquisition/classification/coarse.py or the formal acquisition runner."
  405. )