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
+4 -3
View File
@@ -1,8 +1,9 @@
# md-vector-db 环境变量配置
# 复制此文件为 .env 并填入实际值
# 嵌入 API 密钥 (embed.mode=api 时使用)
EMBED_API_KEY=sk-your-key-here
# 嵌入 API 密钥embed.mode=api 时使用
# 支持的 provider: openai (含硅基流动/智谱/DeepSeek), dashscope (阿里云)
EMBED_API_KEY=your-api-key
# HTTP API 认证密钥 (不设置则跳过认证)
# HTTP API 认证密钥不设置则跳过认证
MD_VECTOR_API_KEY=your-secret-key
+9 -5
View File
@@ -3,14 +3,18 @@ chroma:
collection_name: markdown_docs
embed:
mode: local # local | api
mode: local # local | api
# --- local 模式 ---
local_model: BAAI/bge-small-zh-v1.5
api_base: "" # api 模式下填写
# api_key 请通过环境变量 EMBED_API_KEY 设置
# --- api 模式 ---
provider: openai # openai | dashscope
# api_base: "" # 覆盖默认 API 地址(留空 = provider 默认)
# model: "" # 覆盖默认模型名(留空 = provider 默认)
# api_key 请通过 .env 的 EMBED_API_KEY 设置
chunk:
max_size: 1000 # 分块最大字符数
overlap: 100 # 相邻块重叠字符数
max_size: 1000
overlap: 100
server:
host: 0.0.0.0
+13 -2
View File
@@ -26,11 +26,22 @@ class ChromaConfig:
@dataclass
class EmbedConfig:
"""嵌入模型配置."""
"""嵌入模型配置.
mode: str = "local" # "local" | "api"
Fields:
mode: "local" | "api"
provider: "openai" | "dashscope" (api 模式下的服务商)
local_model: sentence-transformers 模型名 (local 模式)
api_base: 覆盖默认 API 地址 (留空则用 provider 默认值)
api_key: API 密钥 (从 .env 的 EMBED_API_KEY 读取)
model: 覆盖默认模型名 (留空则用 provider 默认值)
"""
mode: str = "local"
provider: str = "openai"
local_model: str = "BAAI/bge-small-zh-v1.5"
api_base: str = ""
model: str = ""
api_key: str = field(
default_factory=lambda: os.environ.get("EMBED_API_KEY", "")
)
+100 -16
View File
@@ -1,12 +1,20 @@
"""嵌入模型抽象层 — 策略模式.
支持的 Provider:
openai — OpenAI / 硅基流动 / 智谱 / DeepSeek / 月之暗面 等 OpenAI 兼容服务
dashscope — 阿里云 DashScope (通义千问)
使用方式:
from src.core.config import EmbedConfig
from src.core.embedder import create_embedder, batch_embed
# 本地模型
embedder = create_embedder(EmbedConfig(mode="local"))
vectors = embedder.embed(["文本1", "文本2"])
# 或分批嵌入避免 OOM:
# OpenAI 兼容 API
embedder = create_embedder(EmbedConfig(mode="api", provider="openai"))
# 阿里云 DashScope
embedder = create_embedder(EmbedConfig(mode="api", provider="dashscope"))
vectors = batch_embed(embedder, long_text_list)
"""
import os
@@ -17,10 +25,24 @@ from src.core.config import EmbedConfig
logger = logging.getLogger(__name__)
# 模型下载镜像 (仅首次下载时使用)
_HF_MIRROR = "https://hf-mirror.com"
# -- Provider 默认配置 --
_PROVIDER_DEFAULTS: dict[str, dict[str, str | int]] = {
"openai": {
"api_base": "https://api.openai.com/v1",
"model": "text-embedding-3-small",
"dimension": 1536,
},
"dashscope": {
"api_base": "https://dashscope.aliyuncs.com/api/v1/services/embeddings/text-embedding/text-embedding",
"model": "text-embedding-v4",
"dimension": 1536,
},
}
# -- 接口 --
class Embedder(Protocol):
"""嵌入器接口."""
@property
@@ -28,8 +50,9 @@ class Embedder(Protocol):
def embed(self, texts: list[str]) -> list[list[float]]: ...
# -- 本地模型 --
class LocalEmbedder:
"""本地 sentence-transformers 模型嵌入器."""
"""sentence-transformers 本地模型嵌入器."""
def __init__(self, config: EmbedConfig):
from sentence_transformers import SentenceTransformer
@@ -61,36 +84,97 @@ class LocalEmbedder:
return embeddings.tolist()
class APIEmbedder:
"""OpenAI 兼容 API 嵌入器."""
# -- API Provider 基类 --
class _BaseAPIEmbedder:
"""API 嵌入器基类: 统一 api_base/model 解析逻辑."""
def __init__(self, config: EmbedConfig):
self._api_base = config.api_base
def __init__(self, config: EmbedConfig, provider: str):
defaults = _PROVIDER_DEFAULTS.get(provider, {})
self._api_base = config.api_base or str(defaults.get("api_base", ""))
self._model = config.model or str(defaults.get("model", ""))
self._api_key = config.api_key
self._dimension = int(defaults.get("dimension", 1536))
@property
def dimension(self) -> int:
return 1536
return self._dimension
def embed(self, texts: list[str]) -> list[list[float]]:
raise NotImplementedError
# -- OpenAI 兼容 API --
class OpenAIEmbedder(_BaseAPIEmbedder):
"""OpenAI / 硅基流动 / 智谱 / DeepSeek / 月之暗面 等."""
def __init__(self, config: EmbedConfig):
super().__init__(config, "openai")
def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
raise ValueError("文本列表不能为空")
from openai import OpenAI
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
response = client.embeddings.create(
model="text-embedding-3-small", input=texts,
)
response = client.embeddings.create(model=self._model, input=texts)
return [d.embedding for d in response.data]
# -- 阿里云 DashScope --
class DashscopeEmbedder(_BaseAPIEmbedder):
"""阿里云 DashScope 嵌入 (自定义 HTTP API, 非 OpenAI 兼容)."""
def __init__(self, config: EmbedConfig):
super().__init__(config, "dashscope")
def embed(self, texts: list[str]) -> list[list[float]]:
if not texts:
raise ValueError("文本列表不能为空")
import requests
resp = requests.post(
self._api_base,
headers={
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
},
json={
"model": self._model,
"input": {"texts": texts},
},
timeout=60,
)
resp.raise_for_status()
data = resp.json()
# DashScope 返回: {"output": {"embeddings": [{"text_index": 0, "embedding": [...]}, ...]}}
embeddings_raw = data.get("output", {}).get("embeddings", [])
# 按 text_index 排序确保顺序
embeddings_raw.sort(key=lambda x: x.get("text_index", 0))
return [e["embedding"] for e in embeddings_raw]
# -- 工厂函数 --
_PROVIDER_CLASSES: dict[str, type[_BaseAPIEmbedder]] = {
"openai": OpenAIEmbedder,
"dashscope": DashscopeEmbedder,
}
SUPPORTED_PROVIDERS = list(_PROVIDER_CLASSES.keys())
def create_embedder(config: EmbedConfig) -> Embedder:
"""工厂函数: 根据配置创建嵌入器."""
if config.mode == "local":
return LocalEmbedder(config)
elif config.mode == "api":
return APIEmbedder(config)
else:
raise ValueError(f"不支持的嵌入模式: {config.mode}")
if config.mode == "api":
provider = config.provider or "openai"
cls = _PROVIDER_CLASSES.get(provider)
if cls is None:
raise ValueError(
f"不支持的 provider: {provider}, 可选: {SUPPORTED_PROVIDERS}"
)
return cls(config)
raise ValueError(f"不支持的嵌入模式: {config.mode}")
def batch_embed(
+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:
"""分批嵌入测试."""