114 lines
4.1 KiB
Python
114 lines
4.1 KiB
Python
"""嵌入模型测试."""
|
|
import pytest
|
|
|
|
from src.core.config import EmbedConfig
|
|
from src.core.embedder import (
|
|
SUPPORTED_PROVIDERS,
|
|
DashscopeEmbedder,
|
|
LocalEmbedder,
|
|
OpenAIEmbedder,
|
|
batch_embed,
|
|
create_embedder,
|
|
)
|
|
|
|
|
|
class TestLocalEmbedder:
|
|
"""本地嵌入器测试."""
|
|
|
|
@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)
|
|
assert isinstance(embedder, LocalEmbedder)
|
|
|
|
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([])
|
|
|
|
|
|
class TestAPIEmbedders:
|
|
"""API 嵌入器测试."""
|
|
|
|
def test_openai_embedder_init(self):
|
|
"""OpenAI 嵌入器使用默认配置."""
|
|
cfg = EmbedConfig(mode="api", provider="openai")
|
|
emb = OpenAIEmbedder(cfg)
|
|
assert emb.dimension == 1536
|
|
assert emb._api_base == "https://api.openai.com/v1"
|
|
|
|
def test_openai_embedder_custom_base(self):
|
|
"""自定义 api_base 覆盖默认值."""
|
|
cfg = EmbedConfig(mode="api", provider="openai",
|
|
api_base="https://api.siliconflow.cn/v1",
|
|
model="BAAI/bge-large-zh-v1.5")
|
|
emb = OpenAIEmbedder(cfg)
|
|
assert emb._api_base == "https://api.siliconflow.cn/v1"
|
|
assert emb._model == "BAAI/bge-large-zh-v1.5"
|
|
|
|
def test_openai_embedder_empty_raises(self):
|
|
"""空列表抛异常."""
|
|
cfg = EmbedConfig(mode="api", provider="openai")
|
|
emb = OpenAIEmbedder(cfg)
|
|
with pytest.raises(ValueError):
|
|
emb.embed([])
|
|
|
|
def test_dashscope_embedder_init(self):
|
|
"""DashScope 嵌入器使用默认配置."""
|
|
cfg = EmbedConfig(mode="api", provider="dashscope")
|
|
emb = DashscopeEmbedder(cfg)
|
|
assert emb.dimension == 1536
|
|
assert emb._model == "text-embedding-v4"
|
|
|
|
def test_factory_creates_correct_provider(self):
|
|
"""工厂函数根据 provider 创建正确的类."""
|
|
openai_emb = create_embedder(EmbedConfig(mode="api", provider="openai"))
|
|
assert isinstance(openai_emb, OpenAIEmbedder)
|
|
ds_emb = create_embedder(EmbedConfig(mode="api", provider="dashscope"))
|
|
assert isinstance(ds_emb, DashscopeEmbedder)
|
|
|
|
def test_factory_rejects_unknown_provider(self):
|
|
"""未知 provider 抛出 ValueError."""
|
|
with pytest.raises(ValueError):
|
|
create_embedder(EmbedConfig(mode="api", provider="unknown"))
|
|
|
|
def test_supported_providers_list(self):
|
|
"""SUPPORTED_PROVIDERS 包含已知 provider."""
|
|
assert "openai" in SUPPORTED_PROVIDERS
|
|
assert "dashscope" in SUPPORTED_PROVIDERS
|
|
|
|
|
|
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
|