diff --git a/src/core/embedder.py b/src/core/embedder.py new file mode 100644 index 0000000..0bcd300 --- /dev/null +++ b/src/core/embedder.py @@ -0,0 +1,59 @@ +"""嵌入模型抽象层.""" +from sentence_transformers import SentenceTransformer + +from src.core.config import EmbedConfig + + +class Embedder: + """文本嵌入器, 支持本地模型和 API 两种模式.""" + + def __init__(self, config: EmbedConfig): + self._config = config + if config.mode == "local": + self._model = SentenceTransformer(config.local_model) + elif config.mode == "api": + self._model = None # 延迟初始化, 需要 openai 包 + self._api_base = config.api_base + self._api_key = config.api_key + else: + raise ValueError(f"不支持的嵌入模式: {config.mode}") + + @property + def mode(self) -> str: + """当前嵌入模式.""" + return self._config.mode + + @property + def dimension(self) -> int: + """嵌入向量维度.""" + if self.mode == "local": + return self._model.get_embedding_dimension() + else: + # 默认 OpenAI text-embedding-ada-002 / text-embedding-3-small 维度 + return 1536 + + def embed(self, texts: list[str]) -> list[list[float]]: + """对文本列表进行嵌入, 返回向量列表.""" + if not texts: + raise ValueError("文本列表不能为空") + if self.mode == "local": + embeddings = self._model.encode(texts, normalize_embeddings=True) + return embeddings.tolist() + else: + return self._embed_via_api(texts) + + def _embed_via_api(self, texts: list[str]) -> list[list[float]]: + """通过 OpenAI 兼容 API 嵌入(延迟导入 openai).""" + from openai import OpenAI + + client = OpenAI(base_url=self._api_base, api_key=self._api_key) + response = client.embeddings.create( + model="text-embedding-3-small", + input=texts, + ) + return [d.embedding for d in response.data] + + +def create_embedder(config: EmbedConfig) -> Embedder: + """工厂函数: 根据配置创建嵌入器.""" + return Embedder(config) diff --git a/tests/test_embedder.py b/tests/test_embedder.py new file mode 100644 index 0000000..1fc6152 --- /dev/null +++ b/tests/test_embedder.py @@ -0,0 +1,51 @@ +"""嵌入模型测试.""" +import pytest + +from src.core.config import EmbedConfig +from src.core.embedder import Embedder, create_embedder + + +class TestEmbedder: + """Embedder 单元测试. + + 注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB), + 首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。 + """ + + @pytest.fixture + def local_config(self): + return EmbedConfig(mode="local") + + def test_create_local_embedder(self, local_config): + """创建本地嵌入器,验证维度正确.""" + embedder = create_embedder(local_config) + assert embedder.dimension > 0 + assert isinstance(embedder.dimension, int) + + def test_embed_single_text(self, local_config): + """嵌入单条文本返回正确维度向量.""" + embedder = create_embedder(local_config) + result = embedder.embed(["你好世界"]) + assert len(result) == 1 + assert len(result[0]) == embedder.dimension + assert all(isinstance(v, float) for v in result[0]) + + def test_embed_multiple_texts(self, local_config): + """嵌入多条文本返回对应数量的向量.""" + embedder = create_embedder(local_config) + texts = ["第一段文本", "第二段文本", "第三段文本"] + result = embedder.embed(texts) + assert len(result) == 3 + for vec in result: + assert len(vec) == embedder.dimension + + def test_embed_empty_list_raises(self, local_config): + """空列表应抛出异常.""" + embedder = create_embedder(local_config) + with pytest.raises(ValueError): + embedder.embed([]) + + def test_create_embedder_from_factory_function(self, local_config): + """工厂函数正确创建 Embedder 实例.""" + embedder = create_embedder(local_config) + assert isinstance(embedder, Embedder)