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 环境变量配置
|
# md-vector-db 环境变量配置
|
||||||
# 复制此文件为 .env 并填入实际值
|
# 复制此文件为 .env 并填入实际值
|
||||||
|
|
||||||
# 嵌入 API 密钥 (embed.mode=api 时使用)
|
# 嵌入 API 密钥(embed.mode=api 时使用)
|
||||||
EMBED_API_KEY=sk-your-key-here
|
# 支持的 provider: openai (含硅基流动/智谱/DeepSeek), dashscope (阿里云)
|
||||||
|
EMBED_API_KEY=your-api-key
|
||||||
|
|
||||||
# HTTP API 认证密钥 (不设置则跳过认证)
|
# HTTP API 认证密钥(不设置则跳过认证)
|
||||||
MD_VECTOR_API_KEY=your-secret-key
|
MD_VECTOR_API_KEY=your-secret-key
|
||||||
|
|||||||
+8
-4
@@ -4,13 +4,17 @@ chroma:
|
|||||||
|
|
||||||
embed:
|
embed:
|
||||||
mode: local # local | api
|
mode: local # local | api
|
||||||
|
# --- local 模式 ---
|
||||||
local_model: BAAI/bge-small-zh-v1.5
|
local_model: BAAI/bge-small-zh-v1.5
|
||||||
api_base: "" # api 模式下填写
|
# --- api 模式 ---
|
||||||
# api_key 请通过环境变量 EMBED_API_KEY 设置
|
provider: openai # openai | dashscope
|
||||||
|
# api_base: "" # 覆盖默认 API 地址(留空 = provider 默认)
|
||||||
|
# model: "" # 覆盖默认模型名(留空 = provider 默认)
|
||||||
|
# api_key 请通过 .env 的 EMBED_API_KEY 设置
|
||||||
|
|
||||||
chunk:
|
chunk:
|
||||||
max_size: 1000 # 分块最大字符数
|
max_size: 1000
|
||||||
overlap: 100 # 相邻块重叠字符数
|
overlap: 100
|
||||||
|
|
||||||
server:
|
server:
|
||||||
host: 0.0.0.0
|
host: 0.0.0.0
|
||||||
|
|||||||
+13
-2
@@ -26,11 +26,22 @@ class ChromaConfig:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EmbedConfig:
|
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"
|
local_model: str = "BAAI/bge-small-zh-v1.5"
|
||||||
api_base: str = ""
|
api_base: str = ""
|
||||||
|
model: str = ""
|
||||||
api_key: str = field(
|
api_key: str = field(
|
||||||
default_factory=lambda: os.environ.get("EMBED_API_KEY", "")
|
default_factory=lambda: os.environ.get("EMBED_API_KEY", "")
|
||||||
)
|
)
|
||||||
|
|||||||
+99
-15
@@ -1,12 +1,20 @@
|
|||||||
"""嵌入模型抽象层 — 策略模式.
|
"""嵌入模型抽象层 — 策略模式.
|
||||||
|
|
||||||
|
支持的 Provider:
|
||||||
|
openai — OpenAI / 硅基流动 / 智谱 / DeepSeek / 月之暗面 等 OpenAI 兼容服务
|
||||||
|
dashscope — 阿里云 DashScope (通义千问)
|
||||||
|
|
||||||
使用方式:
|
使用方式:
|
||||||
from src.core.config import EmbedConfig
|
from src.core.config import EmbedConfig
|
||||||
from src.core.embedder import create_embedder, batch_embed
|
from src.core.embedder import create_embedder, batch_embed
|
||||||
|
|
||||||
|
# 本地模型
|
||||||
embedder = create_embedder(EmbedConfig(mode="local"))
|
embedder = create_embedder(EmbedConfig(mode="local"))
|
||||||
vectors = embedder.embed(["文本1", "文本2"])
|
# OpenAI 兼容 API
|
||||||
# 或分批嵌入避免 OOM:
|
embedder = create_embedder(EmbedConfig(mode="api", provider="openai"))
|
||||||
|
# 阿里云 DashScope
|
||||||
|
embedder = create_embedder(EmbedConfig(mode="api", provider="dashscope"))
|
||||||
|
|
||||||
vectors = batch_embed(embedder, long_text_list)
|
vectors = batch_embed(embedder, long_text_list)
|
||||||
"""
|
"""
|
||||||
import os
|
import os
|
||||||
@@ -17,10 +25,24 @@ from src.core.config import EmbedConfig
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# 模型下载镜像 (仅首次下载时使用)
|
|
||||||
_HF_MIRROR = "https://hf-mirror.com"
|
_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):
|
class Embedder(Protocol):
|
||||||
"""嵌入器接口."""
|
"""嵌入器接口."""
|
||||||
@property
|
@property
|
||||||
@@ -28,8 +50,9 @@ class Embedder(Protocol):
|
|||||||
def embed(self, texts: list[str]) -> list[list[float]]: ...
|
def embed(self, texts: list[str]) -> list[list[float]]: ...
|
||||||
|
|
||||||
|
|
||||||
|
# -- 本地模型 --
|
||||||
class LocalEmbedder:
|
class LocalEmbedder:
|
||||||
"""本地 sentence-transformers 模型嵌入器."""
|
"""sentence-transformers 本地模型嵌入器."""
|
||||||
|
|
||||||
def __init__(self, config: EmbedConfig):
|
def __init__(self, config: EmbedConfig):
|
||||||
from sentence_transformers import SentenceTransformer
|
from sentence_transformers import SentenceTransformer
|
||||||
@@ -61,35 +84,96 @@ class LocalEmbedder:
|
|||||||
return embeddings.tolist()
|
return embeddings.tolist()
|
||||||
|
|
||||||
|
|
||||||
class APIEmbedder:
|
# -- API Provider 基类 --
|
||||||
"""OpenAI 兼容 API 嵌入器."""
|
class _BaseAPIEmbedder:
|
||||||
|
"""API 嵌入器基类: 统一 api_base/model 解析逻辑."""
|
||||||
|
|
||||||
def __init__(self, config: EmbedConfig):
|
def __init__(self, config: EmbedConfig, provider: str):
|
||||||
self._api_base = config.api_base
|
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._api_key = config.api_key
|
||||||
|
self._dimension = int(defaults.get("dimension", 1536))
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def dimension(self) -> int:
|
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]]:
|
def embed(self, texts: list[str]) -> list[list[float]]:
|
||||||
if not texts:
|
if not texts:
|
||||||
raise ValueError("文本列表不能为空")
|
raise ValueError("文本列表不能为空")
|
||||||
from openai import OpenAI
|
from openai import OpenAI
|
||||||
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
|
client = OpenAI(base_url=self._api_base, api_key=self._api_key)
|
||||||
response = client.embeddings.create(
|
response = client.embeddings.create(model=self._model, input=texts)
|
||||||
model="text-embedding-3-small", input=texts,
|
|
||||||
)
|
|
||||||
return [d.embedding for d in response.data]
|
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:
|
def create_embedder(config: EmbedConfig) -> Embedder:
|
||||||
"""工厂函数: 根据配置创建嵌入器."""
|
"""工厂函数: 根据配置创建嵌入器."""
|
||||||
if config.mode == "local":
|
if config.mode == "local":
|
||||||
return LocalEmbedder(config)
|
return LocalEmbedder(config)
|
||||||
elif config.mode == "api":
|
|
||||||
return APIEmbedder(config)
|
if config.mode == "api":
|
||||||
else:
|
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}")
|
raise ValueError(f"不支持的嵌入模式: {config.mode}")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
+47
-10
@@ -2,7 +2,10 @@
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from src.core.config import EmbedConfig
|
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:
|
class TestLocalEmbedder:
|
||||||
@@ -43,22 +46,56 @@ class TestLocalEmbedder:
|
|||||||
embedder.embed([])
|
embedder.embed([])
|
||||||
|
|
||||||
|
|
||||||
class TestAPIEmbedder:
|
class TestAPIEmbedders:
|
||||||
"""API 嵌入器测试."""
|
"""API 嵌入器测试."""
|
||||||
|
|
||||||
def test_api_embedder_init(self):
|
def test_openai_embedder_init(self):
|
||||||
"""API 嵌入器初始化."""
|
"""OpenAI 嵌入器使用默认配置."""
|
||||||
cfg = EmbedConfig(mode="api", api_base="https://api.test.com", api_key="sk-test")
|
cfg = EmbedConfig(mode="api", provider="openai")
|
||||||
emb = APIEmbedder(cfg)
|
emb = OpenAIEmbedder(cfg)
|
||||||
assert emb.dimension == 1536
|
assert emb.dimension == 1536
|
||||||
|
assert emb._api_base == "https://api.openai.com/v1"
|
||||||
|
|
||||||
def test_api_embedder_empty_raises(self):
|
def test_openai_embedder_custom_base(self):
|
||||||
"""API 嵌入器空列表抛异常."""
|
"""自定义 api_base 覆盖默认值."""
|
||||||
cfg = EmbedConfig(mode="api")
|
cfg = EmbedConfig(mode="api", provider="openai",
|
||||||
emb = APIEmbedder(cfg)
|
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):
|
with pytest.raises(ValueError):
|
||||||
emb.embed([])
|
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:
|
class TestBatchEmbed:
|
||||||
"""分批嵌入测试."""
|
"""分批嵌入测试."""
|
||||||
|
|||||||
Reference in New Issue
Block a user