imgtext.py 6.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172
  1. """Image-text multimodal reader for creation knowledge decode."""
  2. from __future__ import annotations
  3. import logging
  4. from typing import Any, Callable, Mapping, Optional
  5. import httpx
  6. from core.config import load_env_file
  7. from core.jsonio import extract_json_object, to_bool
  8. from core.models import Card, CardExtract, ExtractedContent, Post
  9. from core.prompts import load_prompt
  10. logger = logging.getLogger(__name__)
  11. DEFAULT_MODEL = "qwen-vl-plus"
  12. DEFAULT_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
  13. DEFAULT_TIMEOUT = 120.0
  14. MAX_CARDS = 12
  15. def _card_label(card: Card) -> str:
  16. if card.kind == "frame" and card.timestamp is not None:
  17. ts = int(card.timestamp)
  18. return f"【卡片{card.index} · {ts // 60:02d}:{ts % 60:02d}】"
  19. return f"【卡片{card.index}】"
  20. SYSTEM_PROMPT = (
  21. "你是创作知识提取助手。从给定的小红书帖子(标题、正文、图片、视频)中,"
  22. "提取真正能指导『如何创作内容』的知识;知识常在图片/视频里而非正文。"
  23. "忠实提取、不编造。只输出一个 JSON 对象,不要解释或 markdown。"
  24. )
  25. class ExtractorError(RuntimeError):
  26. pass
  27. class BailianExtractor:
  28. def __init__(
  29. self,
  30. *,
  31. api_key: str,
  32. model: str = DEFAULT_MODEL,
  33. base_url: str = DEFAULT_BASE_URL,
  34. timeout_seconds: float = DEFAULT_TIMEOUT,
  35. http_post: Callable[..., Any] = httpx.post,
  36. max_cards: int = MAX_CARDS,
  37. ) -> None:
  38. if not api_key:
  39. raise ExtractorError("missing ALIYUN_BAILIAN_API_KEY")
  40. self.api_key = api_key
  41. self.model = model
  42. self.base_url = base_url.rstrip("/")
  43. self.timeout_seconds = timeout_seconds
  44. self.http_post = http_post
  45. self.max_cards = max_cards
  46. @classmethod
  47. def from_env(cls, env: Mapping[str, str] | None = None, env_file: str = ".env") -> "BailianExtractor":
  48. source = dict(load_env_file(env_file))
  49. if env:
  50. source.update(env)
  51. api_key = source.get("ALIYUN_BAILIAN_API_KEY") or ""
  52. return cls(
  53. api_key=api_key,
  54. model=source.get("ALIYUN_BAILIAN_VL_MODEL") or source.get("ALIYUN_BAILIAN_MODEL") or DEFAULT_MODEL,
  55. base_url=source.get("ALIYUN_BAILIAN_BASE_URL") or DEFAULT_BASE_URL,
  56. timeout_seconds=float(source.get("ALIYUN_BAILIAN_TIMEOUT_SECONDS") or DEFAULT_TIMEOUT),
  57. max_cards=int(source.get("CK_MAX_CARDS") or MAX_CARDS),
  58. )
  59. def _cards(self, post: Post) -> list[Card]:
  60. cards = post.cards or [
  61. Card(index=i, kind="image", url=url)
  62. for i, url in enumerate(post.image_urls, start=1)
  63. ]
  64. if len(cards) > self.max_cards:
  65. dropped = [card.index for card in cards[self.max_cards :]]
  66. logger.warning(
  67. "post %s card count %d exceeds MAX_CARDS=%d, dropping cards %s",
  68. post.id,
  69. len(cards),
  70. self.max_cards,
  71. dropped,
  72. )
  73. cards = cards[: self.max_cards]
  74. return cards
  75. def build_messages(self, post: Post) -> list[dict[str, Any]]:
  76. user_text = load_prompt("extract").format(
  77. title=post.title or "(无)",
  78. topics="、".join(post.topic_list) or "(无)",
  79. body=post.body_text or "(空)",
  80. )
  81. parts: list[dict[str, Any]] = [{"type": "text", "text": user_text}]
  82. for card in self._cards(post):
  83. parts.append({"type": "text", "text": _card_label(card)})
  84. parts.append({"type": "image_url", "image_url": {"url": card.url}})
  85. return [
  86. {"role": "system", "content": SYSTEM_PROMPT},
  87. {"role": "user", "content": parts},
  88. ]
  89. def extract(self, post: Post) -> ExtractedContent:
  90. messages = self.build_messages(post)
  91. last_exc: Optional[Exception] = None
  92. for attempt in range(2):
  93. try:
  94. resp = self.http_post(
  95. f"{self.base_url}/chat/completions",
  96. headers={
  97. "Authorization": f"Bearer {self.api_key}",
  98. "Content-Type": "application/json",
  99. },
  100. json={"model": self.model, "messages": messages},
  101. timeout=self.timeout_seconds,
  102. )
  103. resp.raise_for_status()
  104. content = resp.json()["choices"][0]["message"]["content"]
  105. data = extract_json_object(content)
  106. cards = [
  107. CardExtract(index=int(card["index"]), content=str(card.get("content") or ""))
  108. for card in (data.get("cards") or [])
  109. if isinstance(card, dict) and card.get("index") is not None
  110. ]
  111. return ExtractedContent(
  112. text=str(data.get("text") or ""),
  113. cards=cards,
  114. from_image=str(data.get("from_image") or ""),
  115. from_video=str(data.get("from_video") or ""),
  116. is_empty=to_bool(data.get("is_empty")),
  117. )
  118. except httpx.HTTPError as exc:
  119. last_exc = exc
  120. if attempt == 0:
  121. continue
  122. raise ExtractorError(f"bailian_http_error: {exc}") from exc
  123. except (KeyError, IndexError, TypeError, ValueError) as exc:
  124. last_exc = exc
  125. if attempt == 0:
  126. continue
  127. raise ExtractorError(f"bailian_response_invalid: {exc}") from exc
  128. raise ExtractorError(f"bailian_unknown_error: {last_exc}")
  129. GeminiExtractor = BailianExtractor
  130. def extract_content(
  131. post: Post,
  132. *,
  133. client: Optional[BailianExtractor] = None,
  134. env_file: str = ".env",
  135. ) -> ExtractedContent:
  136. client = client or GeminiExtractor.from_env(env_file=env_file)
  137. return client.extract(post)
  138. def read_imgtext(post: Post, *, extractor: BailianExtractor | None = None) -> ExtractedContent:
  139. client = extractor or BailianExtractor.from_env()
  140. return client.extract(post)
  141. __all__ = [
  142. "BailianExtractor",
  143. "ExtractorError",
  144. "GeminiExtractor",
  145. "extract_content",
  146. "read_imgtext",
  147. ]