From b89c6076e1bb9cd167ad366af1dd4c54b682c8c2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E8=88=AA=E5=AE=87?= <3364451258@qq.com> Date: Sun, 5 Jul 2026 01:43:23 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20multi-provider=20embedder=20=E2=80=94?= =?UTF-8?q?=20support=20OpenAI/DashScope=20with=20provider-specific=20defa?= =?UTF-8?q?ults?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .env.example | 7 +-- config.yaml | 14 +++-- src/core/config.py | 15 +++++- src/core/embedder.py | 116 +++++++++++++++++++++++++++++++++++------ tests/test_embedder.py | 57 ++++++++++++++++---- 5 files changed, 173 insertions(+), 36 deletions(-) diff --git a/.env.example b/.env.example index 08041c9..7365c30 100644 --- a/.env.example +++ b/.env.example @@ -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 diff --git a/config.yaml b/config.yaml index ec1ada4..bf9bdb7 100644 --- a/config.yaml +++ b/config.yaml @@ -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 diff --git a/src/core/config.py b/src/core/config.py index ff7cd99..10d8228 100644 --- a/src/core/config.py +++ b/src/core/config.py @@ -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", "") ) diff --git a/src/core/embedder.py b/src/core/embedder.py index 437a1f8..1132a62 100644 --- a/src/core/embedder.py +++ b/src/core/embedder.py @@ -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( diff --git a/tests/test_embedder.py b/tests/test_embedder.py index da9ecbf..a4d9a43 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -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: """分批嵌入测试."""