Files
2026-07-18 15:51:01 +08:00

184 lines
6.2 KiB
Python

"""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()