feat: multi-provider embedder — support OpenAI/DashScope with provider-specific defaults

This commit is contained in:
2026-07-05 01:43:23 +08:00
parent 5672472030
commit b89c6076e1
5 changed files with 173 additions and 36 deletions
+47 -10
View File
@@ -2,7 +2,10 @@
import pytest
from src.core.config import EmbedConfig
from src.core.embedder import LocalEmbedder, APIEmbedder, create_embedder, batch_embed
from src.core.embedder import (
LocalEmbedder, OpenAIEmbedder, DashscopeEmbedder,
create_embedder, batch_embed, SUPPORTED_PROVIDERS,
)
class TestLocalEmbedder:
@@ -43,22 +46,56 @@ class TestLocalEmbedder:
embedder.embed([])
class TestAPIEmbedder:
class TestAPIEmbedders:
"""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)
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_api_embedder_empty_raises(self):
"""API 嵌入器空列表抛异常."""
cfg = EmbedConfig(mode="api")
emb = APIEmbedder(cfg)
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:
"""分批嵌入测试."""