imgtext.py 9.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245
  1. """Image-text multimodal reader for creation knowledge decode."""
  2. from __future__ import annotations
  3. import logging
  4. import time
  5. from typing import Any, Callable, Mapping, Optional
  6. import httpx
  7. from core.config import load_env_file
  8. from core.jsonio import extract_json_object, to_bool
  9. from core.models import Card, CardExtract, ExtractedContent, Post
  10. from core.prompts import load_prompt
  11. from pipeline.tracing import TraceContext, TraceWriter, hash_prompt, redact_headers, timed_ms
  12. logger = logging.getLogger(__name__)
  13. DEFAULT_MODEL = "qwen-vl-plus"
  14. DEFAULT_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
  15. DEFAULT_TIMEOUT = 120.0
  16. MAX_CARDS = 12
  17. def _card_label(card: Card) -> str:
  18. if card.kind == "frame" and card.timestamp is not None:
  19. ts = int(card.timestamp)
  20. return f"【卡片{card.index} · {ts // 60:02d}:{ts % 60:02d}】"
  21. return f"【卡片{card.index}】"
  22. SYSTEM_PROMPT = (
  23. "你是创作知识提取助手。从给定的小红书帖子(标题、正文、图片、视频)中,"
  24. "提取真正能指导『如何创作内容』的知识;知识常在图片/视频里而非正文。"
  25. "忠实提取、不编造。只输出一个 JSON 对象,不要解释或 markdown。"
  26. )
  27. class ExtractorError(RuntimeError):
  28. pass
  29. class BailianExtractor:
  30. def __init__(
  31. self,
  32. *,
  33. api_key: str,
  34. model: str = DEFAULT_MODEL,
  35. base_url: str = DEFAULT_BASE_URL,
  36. timeout_seconds: float = DEFAULT_TIMEOUT,
  37. http_post: Callable[..., Any] = httpx.post,
  38. max_cards: int = MAX_CARDS,
  39. ) -> None:
  40. if not api_key:
  41. raise ExtractorError("missing ALIYUN_BAILIAN_API_KEY")
  42. self.api_key = api_key
  43. self.model = model
  44. self.base_url = base_url.rstrip("/")
  45. self.timeout_seconds = timeout_seconds
  46. self.http_post = http_post
  47. self.max_cards = max_cards
  48. @classmethod
  49. def from_env(cls, env: Mapping[str, str] | None = None, env_file: str = ".env") -> "BailianExtractor":
  50. source = dict(load_env_file(env_file))
  51. if env:
  52. source.update(env)
  53. api_key = source.get("ALIYUN_BAILIAN_API_KEY") or ""
  54. return cls(
  55. api_key=api_key,
  56. model=source.get("ALIYUN_BAILIAN_VL_MODEL") or source.get("ALIYUN_BAILIAN_MODEL") or DEFAULT_MODEL,
  57. base_url=source.get("ALIYUN_BAILIAN_BASE_URL") or DEFAULT_BASE_URL,
  58. timeout_seconds=float(source.get("ALIYUN_BAILIAN_TIMEOUT_SECONDS") or DEFAULT_TIMEOUT),
  59. max_cards=int(source.get("CK_MAX_CARDS") or MAX_CARDS),
  60. )
  61. def _cards(self, post: Post) -> list[Card]:
  62. cards = post.cards or [
  63. Card(index=i, kind="image", url=url)
  64. for i, url in enumerate(post.image_urls, start=1)
  65. ]
  66. if len(cards) > self.max_cards:
  67. dropped = [card.index for card in cards[self.max_cards :]]
  68. logger.warning(
  69. "post %s card count %d exceeds MAX_CARDS=%d, dropping cards %s",
  70. post.id,
  71. len(cards),
  72. self.max_cards,
  73. dropped,
  74. )
  75. cards = cards[: self.max_cards]
  76. return cards
  77. def build_messages(self, post: Post) -> list[dict[str, Any]]:
  78. user_text = load_prompt("extract").format(
  79. title=post.title or "(无)",
  80. topics="、".join(post.topic_list) or "(无)",
  81. body=post.body_text or "(空)",
  82. )
  83. parts: list[dict[str, Any]] = [{"type": "text", "text": user_text}]
  84. for card in self._cards(post):
  85. parts.append({"type": "text", "text": _card_label(card)})
  86. parts.append({"type": "image_url", "image_url": {"url": card.url}})
  87. return [
  88. {"role": "system", "content": SYSTEM_PROMPT},
  89. {"role": "user", "content": parts},
  90. ]
  91. def extract(
  92. self,
  93. post: Post,
  94. *,
  95. trace_writer: TraceWriter | None = None,
  96. trace_context: TraceContext | None = None,
  97. ) -> ExtractedContent:
  98. messages = self.build_messages(post)
  99. last_exc: Optional[Exception] = None
  100. for attempt in range(2):
  101. started = time.perf_counter()
  102. headers = {
  103. "Authorization": f"Bearer {self.api_key}",
  104. "Content-Type": "application/json",
  105. }
  106. body = {"model": self.model, "messages": messages}
  107. try:
  108. resp = self.http_post(
  109. f"{self.base_url}/chat/completions",
  110. headers=headers,
  111. json=body,
  112. timeout=self.timeout_seconds,
  113. )
  114. resp.raise_for_status()
  115. response_json = resp.json()
  116. content = response_json["choices"][0]["message"]["content"]
  117. data = extract_json_object(content)
  118. cards = [
  119. CardExtract(index=int(card["index"]), content=str(card.get("content") or ""))
  120. for card in (data.get("cards") or [])
  121. if isinstance(card, dict) and card.get("index") is not None
  122. ]
  123. if trace_writer is not None:
  124. trace_writer.llm_call(
  125. context=trace_context or TraceContext(stage="decode", substage="read_imgtext"),
  126. stage="decode",
  127. substage="read_imgtext",
  128. provider="bailian",
  129. model_name=self.model,
  130. endpoint=f"{self.base_url}/chat/completions",
  131. prompt_name="extract_imgtext",
  132. prompt_hash=hash_prompt(SYSTEM_PROMPT),
  133. request_payload={"headers": redact_headers(headers), **body},
  134. response_payload=response_json,
  135. parsed_payload={
  136. "is_empty": to_bool(data.get("is_empty")),
  137. "card_count": len(cards),
  138. "text": data.get("text"),
  139. },
  140. status="done",
  141. latency_ms=timed_ms(started),
  142. attempt_index=attempt + 1,
  143. )
  144. return ExtractedContent(
  145. text=str(data.get("text") or ""),
  146. cards=cards,
  147. from_image=str(data.get("from_image") or ""),
  148. from_video=str(data.get("from_video") or ""),
  149. is_empty=to_bool(data.get("is_empty")),
  150. )
  151. except httpx.HTTPError as exc:
  152. last_exc = exc
  153. if trace_writer is not None:
  154. trace_writer.llm_call(
  155. context=trace_context or TraceContext(stage="decode", substage="read_imgtext"),
  156. stage="decode",
  157. substage="read_imgtext",
  158. provider="bailian",
  159. model_name=self.model,
  160. endpoint=f"{self.base_url}/chat/completions",
  161. prompt_name="extract_imgtext",
  162. prompt_hash=hash_prompt(SYSTEM_PROMPT),
  163. request_payload={"headers": redact_headers(headers), **body},
  164. status="failed",
  165. error_message=str(exc),
  166. latency_ms=timed_ms(started),
  167. attempt_index=attempt + 1,
  168. )
  169. if attempt == 0:
  170. continue
  171. raise ExtractorError(f"bailian_http_error: {exc}") from exc
  172. except (KeyError, IndexError, TypeError, ValueError) as exc:
  173. last_exc = exc
  174. if trace_writer is not None:
  175. trace_writer.llm_call(
  176. context=trace_context or TraceContext(stage="decode", substage="read_imgtext"),
  177. stage="decode",
  178. substage="read_imgtext",
  179. provider="bailian",
  180. model_name=self.model,
  181. endpoint=f"{self.base_url}/chat/completions",
  182. prompt_name="extract_imgtext",
  183. prompt_hash=hash_prompt(SYSTEM_PROMPT),
  184. request_payload={"headers": redact_headers(headers), **body},
  185. status="failed",
  186. error_message=str(exc),
  187. latency_ms=timed_ms(started),
  188. attempt_index=attempt + 1,
  189. )
  190. if attempt == 0:
  191. continue
  192. raise ExtractorError(f"bailian_response_invalid: {exc}") from exc
  193. raise ExtractorError(f"bailian_unknown_error: {last_exc}")
  194. GeminiExtractor = BailianExtractor
  195. def extract_content(
  196. post: Post,
  197. *,
  198. client: Optional[BailianExtractor] = None,
  199. env_file: str = ".env",
  200. trace_writer: TraceWriter | None = None,
  201. trace_context: TraceContext | None = None,
  202. ) -> ExtractedContent:
  203. client = client or GeminiExtractor.from_env(env_file=env_file)
  204. return client.extract(post, trace_writer=trace_writer, trace_context=trace_context)
  205. def read_imgtext(
  206. post: Post,
  207. *,
  208. extractor: BailianExtractor | None = None,
  209. trace_writer: TraceWriter | None = None,
  210. trace_context: TraceContext | None = None,
  211. ) -> ExtractedContent:
  212. client = extractor or BailianExtractor.from_env()
  213. return client.extract(post, trace_writer=trace_writer, trace_context=trace_context)
  214. __all__ = [
  215. "BailianExtractor",
  216. "ExtractorError",
  217. "GeminiExtractor",
  218. "extract_content",
  219. "read_imgtext",
  220. ]