Initial commit
This commit is contained in:
@@ -0,0 +1,183 @@
|
||||
"""DashScope / Qwen 远程嵌入实现。
|
||||
|
||||
通过 OpenAI 兼容接口调用阿里百炼的 text-embedding-v3:
|
||||
base_url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||
model: text-embedding-v3 (1024 维)
|
||||
限制: 单次请求 input ≤ 25 条
|
||||
|
||||
环境变量:
|
||||
DASHSCOPE_API_KEY
|
||||
QWEN_BASE_URL (默认百炼兼容路径)
|
||||
EMBEDDING_MODEL (默认 text-embedding-v3)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
|
||||
from loguru import logger
|
||||
from openai import AsyncOpenAI, OpenAI
|
||||
|
||||
from .base import AsyncEmbeddingProvider, EmbeddingProvider
|
||||
from .models import EmbeddingError
|
||||
|
||||
DASHSCOPE_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
DASHSCOPE_DEFAULT_MODEL = "text-embedding-v3"
|
||||
DASHSCOPE_DEFAULT_DIM = 1024
|
||||
DASHSCOPE_BATCH_LIMIT = 10 # 百炼实测单批上限(2026-06,文档曾标 25 但 API 报 400)
|
||||
|
||||
# 重试策略
|
||||
DEFAULT_MAX_ATTEMPTS = 3
|
||||
RETRY_BASE_WAIT_SEC = 1.0
|
||||
RETRY_MAX_WAIT_SEC = 8.0
|
||||
|
||||
|
||||
def _read_env(key: str, default: str | None = None) -> str | None:
|
||||
val = os.environ.get(key)
|
||||
if val is None or val.strip() == "":
|
||||
return default
|
||||
return val.strip()
|
||||
|
||||
|
||||
def _resolve_config() -> tuple[str, str, str]:
|
||||
"""读取 API key / base_url / model,返回 (api_key, base_url, model)。
|
||||
|
||||
模型名优先级:DASHSCOPE_EMBEDDING_MODEL > 默认值。
|
||||
不再读全局 EMBEDDING_MODEL,避免与 LOCAL provider 冲突。
|
||||
"""
|
||||
api_key = _read_env("DASHSCOPE_EMBEDDING_API_KEY") or _read_env("DASHSCOPE_API_KEY") or ""
|
||||
if not api_key:
|
||||
raise EmbeddingError("DASHSCOPE_EMBEDDING_API_KEY 或 DASHSCOPE_API_KEY 未配置")
|
||||
# DASHSCOPE_EMBEDDING_BASE_URL -> QWEN_BASE_URL(兜底) -> 默认
|
||||
base_url = (
|
||||
_read_env("DASHSCOPE_EMBEDDING_BASE_URL")
|
||||
or _read_env("QWEN_BASE_URL")
|
||||
or DASHSCOPE_DEFAULT_BASE
|
||||
)
|
||||
model = (
|
||||
_read_env("DASHSCOPE_EMBEDDING_MODEL", DASHSCOPE_DEFAULT_MODEL)
|
||||
or DASHSCOPE_DEFAULT_MODEL
|
||||
)
|
||||
return api_key, base_url, model
|
||||
|
||||
|
||||
def _chunked(items: list[str], size: int) -> list[list[str]]:
|
||||
"""把列表按 size 分块。"""
|
||||
return [items[i : i + size] for i in range(0, len(items), size)]
|
||||
|
||||
|
||||
class DashScopeEmbeddingProvider(EmbeddingProvider):
|
||||
"""同步实现,主要用于测试/单条调用。"""
|
||||
|
||||
name = "dashscope"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
timeout_sec: float = 60.0,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
) -> None:
|
||||
env_key, env_base, env_model = _resolve_config()
|
||||
self.model = model or env_model
|
||||
self.dim = DASHSCOPE_DEFAULT_DIM
|
||||
self.max_attempts = max_attempts
|
||||
self._client = OpenAI(
|
||||
api_key=api_key or env_key,
|
||||
base_url=base_url or env_base,
|
||||
timeout=timeout_sec,
|
||||
)
|
||||
|
||||
def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
results: list[list[float]] = []
|
||||
for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT):
|
||||
results.extend(self._call_with_retry(chunk))
|
||||
return results
|
||||
|
||||
def _call_with_retry(self, batch: list[str]) -> list[list[float]]:
|
||||
import time
|
||||
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(1, self.max_attempts + 1):
|
||||
try:
|
||||
resp = self._client.embeddings.create(model=self.model, input=batch)
|
||||
return [d.embedding for d in resp.data]
|
||||
except Exception as e: # noqa: BLE001
|
||||
last_err = e
|
||||
logger.warning(
|
||||
"DashScope embed 失败 尝试 {}/{}: {}: {}",
|
||||
attempt, self.max_attempts, type(e).__name__, e,
|
||||
)
|
||||
if attempt < self.max_attempts:
|
||||
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
|
||||
time.sleep(wait)
|
||||
raise EmbeddingError(
|
||||
f"DashScope embed 放弃 {self.max_attempts} 次: {last_err}",
|
||||
attempts=self.max_attempts,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self._client.close()
|
||||
|
||||
|
||||
class DashScopeAsyncEmbeddingProvider(AsyncEmbeddingProvider):
|
||||
"""异步实现,用于批处理。"""
|
||||
|
||||
name = "dashscope"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model: str | None = None,
|
||||
api_key: str | None = None,
|
||||
base_url: str | None = None,
|
||||
timeout_sec: float = 60.0,
|
||||
max_attempts: int = DEFAULT_MAX_ATTEMPTS,
|
||||
) -> None:
|
||||
env_key, env_base, env_model = _resolve_config()
|
||||
self.model = model or env_model
|
||||
self.dim = DASHSCOPE_DEFAULT_DIM
|
||||
self.max_attempts = max_attempts
|
||||
self._client = AsyncOpenAI(
|
||||
api_key=api_key or env_key,
|
||||
base_url=base_url or env_base,
|
||||
timeout=timeout_sec,
|
||||
)
|
||||
|
||||
async def embed_batch(self, texts: list[str]) -> list[list[float]]:
|
||||
if not texts:
|
||||
return []
|
||||
results: list[list[float]] = []
|
||||
for chunk in _chunked(texts, DASHSCOPE_BATCH_LIMIT):
|
||||
results.extend(await self._call_with_retry(chunk))
|
||||
return results
|
||||
|
||||
async def _call_with_retry(self, batch: list[str]) -> list[list[float]]:
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(1, self.max_attempts + 1):
|
||||
try:
|
||||
resp = await self._client.embeddings.create(
|
||||
model=self.model, input=batch
|
||||
)
|
||||
return [d.embedding for d in resp.data]
|
||||
except Exception as e: # noqa: BLE001
|
||||
last_err = e
|
||||
logger.warning(
|
||||
"DashScope embed 失败 尝试 {}/{}: {}: {}",
|
||||
attempt, self.max_attempts, type(e).__name__, e,
|
||||
)
|
||||
if attempt < self.max_attempts:
|
||||
wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC)
|
||||
await asyncio.sleep(wait)
|
||||
raise EmbeddingError(
|
||||
f"DashScope embed 放弃 {self.max_attempts} 次: {last_err}",
|
||||
attempts=self.max_attempts,
|
||||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
await self._client.close()
|
||||
Reference in New Issue
Block a user