"""嵌入模型测试.""" import pytest from src.core.config import EmbedConfig from src.core.embedder import ( LocalEmbedder, OpenAIEmbedder, DashscopeEmbedder, create_embedder, batch_embed, SUPPORTED_PROVIDERS, ) 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