feat: multi-provider embedder — support OpenAI/DashScope with provider-specific defaults
This commit is contained in:
+47
-10
@@ -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:
|
||||
"""分批嵌入测试."""
|
||||
|
||||
Reference in New Issue
Block a user