refactor: strategy-pattern embedder (Local/API), remove global env vars, add batch embedding

This commit is contained in:
2026-07-05 01:35:54 +08:00
parent 81a8173a38
commit ad9e5ce449
3 changed files with 112 additions and 66 deletions
+32 -11
View File
@@ -2,15 +2,11 @@
import pytest
from src.core.config import EmbedConfig
from src.core.embedder import Embedder, create_embedder
from src.core.embedder import LocalEmbedder, APIEmbedder, create_embedder, batch_embed
class TestEmbedder:
"""Embedder 单元测试.
注意:本地模型测试需要下载 sentence-transformers 模型(约 100MB),
首次运行耗时较长。API 模式测试使用 mock 避免网络依赖。
"""
class TestLocalEmbedder:
"""本地嵌入器测试."""
@pytest.fixture
def local_config(self):
@@ -21,6 +17,7 @@ class TestEmbedder:
embedder = create_embedder(local_config)
assert embedder.dimension > 0
assert isinstance(embedder.dimension, int)
assert isinstance(embedder, LocalEmbedder)
def test_embed_single_text(self, local_config):
"""嵌入单条文本返回正确维度向量."""
@@ -45,7 +42,31 @@ class TestEmbedder:
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)
class TestAPIEmbedder:
"""API 嵌入器测试."""
def test_api_embedder_init(self):
"""API 嵌入器初始化."""
cfg = EmbedConfig(mode="api", api_base="https://api.test.com", api_key="sk-test")
emb = APIEmbedder(cfg)
assert emb.dimension == 1536
def test_api_embedder_empty_raises(self):
"""API 嵌入器空列表抛异常."""
cfg = EmbedConfig(mode="api")
emb = APIEmbedder(cfg)
with pytest.raises(ValueError):
emb.embed([])
class TestBatchEmbed:
"""分批嵌入测试."""
def test_batch_embed(self):
"""分批嵌入返回正确数量."""
embedder = create_embedder(EmbedConfig(mode="local"))
texts = ["测试文本"] * 70 # 超过 batch_size 32
results = batch_embed(embedder, texts, batch_size=32)
assert len(results) == 70
assert len(results[0]) == embedder.dimension