coarse.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199
  1. """Formal coarse classifier for creation-knowledge candidates."""
  2. from __future__ import annotations
  3. import hashlib
  4. from dataclasses import dataclass, field
  5. from pathlib import Path
  6. from typing import Any
  7. from acquisition.classify import (
  8. MAX_CARDS,
  9. _data_url,
  10. _is_http_url,
  11. _judge,
  12. classify_video as _legacy_classify_video,
  13. )
  14. from core.config import Settings
  15. from core.prompts import load_prompt
  16. from core.text_limits import CLASSIFY_BODY_MAX_CHARS, clip_text
  17. from pipeline.tracing import TraceContext, TraceWriter
  18. ROOT = Path(__file__).resolve().parents[2]
  19. @dataclass(frozen=True)
  20. class ClassificationResult:
  21. is_creation_knowledge: bool | None
  22. label: str | None
  23. confidence: float | None
  24. reason: str
  25. knowledge: str = ""
  26. prompt_version: str | None = None
  27. result_payload: dict[str, Any] = field(default_factory=dict)
  28. status: str = "classified"
  29. error_message: str | None = None
  30. def prompt_version(*names: str) -> str:
  31. h = hashlib.sha256()
  32. for name in names:
  33. path = ROOT / "prompts" / f"{name}.txt"
  34. h.update(name.encode("utf-8"))
  35. if path.exists():
  36. h.update(path.read_bytes())
  37. return h.hexdigest()[:16]
  38. def classify_imgtext(
  39. payload: dict[str, Any],
  40. settings: Settings,
  41. *,
  42. trace_writer: TraceWriter | None = None,
  43. trace_context: TraceContext | None = None,
  44. ) -> tuple:
  45. """Classify image-text content, accepting both HTTP image URLs and /data paths."""
  46. user = [
  47. {
  48. "type": "text",
  49. "text": (
  50. f"平台:{payload.get('platform')}\n"
  51. f"标题:{payload.get('title', '')}\n"
  52. f"正文:{clip_text(payload.get('body_text') or '', CLASSIFY_BODY_MAX_CHARS)}\n"
  53. "(下附帖子图片,请一并看完)"
  54. ),
  55. }
  56. ]
  57. for image in (payload.get("images") or [])[:MAX_CARDS]:
  58. if _is_http_url(image):
  59. user.append({"type": "image_url", "image_url": {"url": image}})
  60. continue
  61. data_url = _data_url(image, settings)
  62. if data_url:
  63. user.append({"type": "image_url", "image_url": {"url": data_url}})
  64. messages = [
  65. {"role": "system", "content": load_prompt("classify_imgtext")},
  66. {"role": "user", "content": user},
  67. ]
  68. if trace_writer is None and trace_context is None:
  69. return _judge(messages, settings, timeout=120)
  70. return _judge(
  71. messages,
  72. settings,
  73. timeout=120,
  74. trace_writer=trace_writer,
  75. trace_context=trace_context,
  76. trace_stage="classify",
  77. trace_substage="coarse_imgtext",
  78. prompt_name="classify_imgtext",
  79. )
  80. def classify_video(
  81. payload: dict[str, Any],
  82. settings: Settings,
  83. *,
  84. trace_writer: TraceWriter | None = None,
  85. trace_context: TraceContext | None = None,
  86. ) -> tuple:
  87. if trace_writer is None and trace_context is None:
  88. return _legacy_classify_video(payload, settings)
  89. return _legacy_classify_video(
  90. payload,
  91. settings,
  92. trace_writer=trace_writer,
  93. trace_context=trace_context,
  94. )
  95. def coarse_classify_item(
  96. *,
  97. platform: str,
  98. content_mode: str | None = None,
  99. title: str = "",
  100. body_text: str = "",
  101. image_urls: list[str] | None = None,
  102. video_url: str = "",
  103. settings: Settings,
  104. trace_writer: TraceWriter | None = None,
  105. trace_context: TraceContext | None = None,
  106. ) -> ClassificationResult:
  107. if content_mode == "unsupported":
  108. return ClassificationResult(
  109. is_creation_knowledge=None,
  110. label="unsupported_content_mode",
  111. confidence=None,
  112. reason="内容模态暂不支持,跳过粗筛",
  113. prompt_version="content_mode_guard",
  114. result_payload={"content_mode": content_mode},
  115. status="skipped",
  116. error_message="unsupported_content_mode",
  117. )
  118. if content_mode == "video_post" and not video_url:
  119. return ClassificationResult(
  120. is_creation_knowledge=None,
  121. label="video_missing",
  122. confidence=None,
  123. reason="视频帖缺少可处理的视频地址,跳过粗筛",
  124. prompt_version="content_mode_guard",
  125. result_payload={"content_mode": content_mode, "video_url_missing": True},
  126. status="skipped",
  127. error_message="video_url_missing",
  128. )
  129. if content_mode == "video_post" or (content_mode is None and video_url):
  130. version = prompt_version("classify_video")
  131. payload = {
  132. "platform": platform,
  133. "title": title,
  134. "body_text": body_text,
  135. "video": video_url,
  136. }
  137. if trace_writer is None and trace_context is None:
  138. is_creation, reason, knowledge, points = classify_video(payload, settings)
  139. else:
  140. is_creation, reason, knowledge, points = classify_video(
  141. payload,
  142. settings,
  143. trace_writer=trace_writer,
  144. trace_context=trace_context,
  145. )
  146. else:
  147. version = prompt_version("classify_imgtext")
  148. payload = {
  149. "platform": platform,
  150. "title": title,
  151. "body_text": body_text,
  152. "images": image_urls or [],
  153. }
  154. if trace_writer is None and trace_context is None:
  155. is_creation, reason, knowledge, points = classify_imgtext(payload, settings)
  156. else:
  157. is_creation, reason, knowledge, points = classify_imgtext(
  158. payload,
  159. settings,
  160. trace_writer=trace_writer,
  161. trace_context=trace_context,
  162. )
  163. if is_creation is None:
  164. return ClassificationResult(
  165. is_creation_knowledge=None,
  166. label=None,
  167. confidence=None,
  168. reason=reason,
  169. knowledge=knowledge,
  170. prompt_version=version,
  171. result_payload={"knowledge": knowledge, "points": points},
  172. status="failed",
  173. error_message=reason,
  174. )
  175. is_hit = bool(is_creation)
  176. return ClassificationResult(
  177. is_creation_knowledge=is_hit,
  178. label="creation" if is_hit else "not_creation",
  179. confidence=1.0,
  180. reason=reason,
  181. knowledge=knowledge,
  182. prompt_version=version,
  183. result_payload={"knowledge": knowledge, "points": points},
  184. status="classified",
  185. )