feat: multi-provider embedder — support OpenAI/DashScope with provider-specific defaults
This commit is contained in:
+4
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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