"""Embedding provider 抽象接口与文本组装工具。 文本组装策略 (compose_text): 优先组装 ExtractedEvent 时,把"语义浓缩"信息前置: title | sentiment+importance+event_type | summary | content[截断] 退化为 Article 时: title | content[截断] 超长截断保护:默认 4000 字符(BGE-M3 max_seq=8192,远程也保守取值)。 """ from __future__ import annotations from abc import ABC, abstractmethod from datetime import datetime from typing import Any from extractor import Article # 拼接后送入 embedder 的字符数上限 MAX_TEXT_CHARS = 4000 def _from_event_dict(event_obj: dict[str, Any]) -> tuple[Article, str, str | None]: """从 ExtractedEvent JSON dict 中,取出 Article 元数据 + 加权 head 段。 返回 (article, head, summary): head 为"事件标签摘要",会作为前缀拼到嵌入文本前; summary 为 event.summary。 """ title = event_obj.get("title") or "" url = event_obj.get("url") or "" url_hash = event_obj.get("url_hash") or "" source_id = event_obj.get("source_id") or "" publish_time_raw = event_obj.get("publish_time") publish_time = ( datetime.fromisoformat(publish_time_raw) if isinstance(publish_time_raw, str) and publish_time_raw else None ) ev = event_obj.get("event") or {} sentiment = ev.get("sentiment") or "neutral" importance = ev.get("importance") or 0 event_type = ev.get("event_type") or "其他" stock_codes = ev.get("stock_codes") or [] company_names = ev.get("company_names") or [] industries = ev.get("industries") or [] summary = ev.get("summary") or "" head_parts = [ f"sentiment={sentiment}", f"importance={importance}", f"event_type={event_type}", ] if company_names: head_parts.append("公司=" + ",".join(company_names[:5])) if industries: head_parts.append("行业=" + ",".join(industries[:5])) if stock_codes: head_parts.append("代码=" + ",".join(stock_codes[:5])) head = "[" + " ".join(head_parts) + "]" # 用 ExtractedEvent 中存在的字段构造一个最小 Article 让下游兼容 article = Article( source_id=source_id, url=url, url_hash=url_hash, title=title, content=ev.get("summary") or title, # 占位,真正的正文从原 article 文件读 publish_time=publish_time, word_count=0, ) return article, head, summary def compose_text( article: Article, *, head: str | None = None, summary: str | None = None, max_chars: int = MAX_TEXT_CHARS, ) -> str: """把 Article 组装成单段嵌入文本。 参数: article: 输入文章(用于 title + content) head: 可选事件标签摘要(由 ExtractedEvent 提取),前置可提高检索信号 summary: 可选 LLM 生成的一句话摘要,前置 head 之后 max_chars: 整段最大字符数,超出截断 content """ parts: list[str] = [f"标题:{article.title}"] if head: parts.append(head) if summary: parts.append(f"摘要:{summary}") body = article.content or "" parts.append("正文:" + body) text = "\n".join(parts) if len(text) > max_chars: text = text[:max_chars] return text # --------------------------------------------------------------------------- # # Provider 抽象接口 # --------------------------------------------------------------------------- # class EmbeddingProvider(ABC): """同步嵌入 provider 抽象。""" name: str model: str dim: int @abstractmethod def embed_batch(self, texts: list[str]) -> list[list[float]]: """批量嵌入,返回与输入等长的向量列表。""" def embed_one(self, text: str) -> list[float]: """单条嵌入,默认走 batch=1。""" return self.embed_batch([text])[0] def close(self) -> None: # noqa: B027 - 默认空实现,子类按需覆盖 """释放资源(如 HTTP client / 模型)。""" def __enter__(self) -> EmbeddingProvider: return self def __exit__(self, *_: object) -> None: self.close() class AsyncEmbeddingProvider(ABC): """异步嵌入 provider(用于批处理高并发)。""" name: str model: str dim: int @abstractmethod async def embed_batch(self, texts: list[str]) -> list[list[float]]: ... async def embed_one(self, text: str) -> list[float]: return (await self.embed_batch([text]))[0] async def close(self) -> None: # noqa: B027 - 默认空实现 ... async def __aenter__(self) -> AsyncEmbeddingProvider: return self async def __aexit__(self, *_: object) -> None: await self.close()