"""DashScope Embedding 客户端。 通过 OpenAI 兼容接口调用阿里百炼 text-embedding-v3: base_url: https://dashscope.aliyuncs.com/compatible-mode/v1 model: text-embedding-v3(1024 维) 配置来源: - .env → DASHSCOPE_API_KEY / QWEN_BASE_URL(OpenAI 兼容端点) - configs/system.yaml → embedding 段(model / dimension / batch_size / timeout) """ import logging import os import time from dataclasses import dataclass from pathlib import Path import yaml from openai import OpenAI from embedding.models import EmbeddingError logger = logging.getLogger(__name__) # 默认值 DASHSCOPE_DEFAULT_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" DASHSCOPE_DEFAULT_MODEL = "text-embedding-v3" DASHSCOPE_DEFAULT_DIM = 1024 DASHSCOPE_BATCH_LIMIT = 10 # 百炼实测单批上限 # 重试 DEFAULT_MAX_ATTEMPTS = 3 RETRY_BASE_WAIT_SEC = 1.0 RETRY_MAX_WAIT_SEC = 8.0 def _load_embedding_config() -> dict: """从 system.yaml 加载 embedding 段配置。""" config_path = Path("configs/system.yaml") if config_path.exists(): try: with open(config_path, encoding="utf-8") as f: raw = yaml.safe_load(f) return raw.get("embedding", {}) except Exception: logger.warning("加载 embedding 配置失败") return {} @dataclass class EmbeddingConfig: """Embedding 调用配置。""" provider: str = "dashscope" model: str = DASHSCOPE_DEFAULT_MODEL api_key: str = "" base_url: str = DASHSCOPE_DEFAULT_BASE dimension: int = DASHSCOPE_DEFAULT_DIM batch_size: int = DASHSCOPE_BATCH_LIMIT timeout_sec: float = 30.0 max_attempts: int = DEFAULT_MAX_ATTEMPTS def __post_init__(self) -> None: if not self.api_key: raise EmbeddingError("DASHSCOPE_API_KEY 未配置,请检查 .env") def load_embedding_config( *, model: str | None = None, ) -> EmbeddingConfig: """根据配置构造 EmbeddingConfig。 优先级: system.yaml > .env 默认值 > 硬编码默认值 """ sys_cfg = _load_embedding_config() # API key: 从环境变量读取 api_key = os.environ.get("DASHSCOPE_API_KEY", "") if not api_key: api_key = os.environ.get("QWEN_API_KEY", "") # Base URL: 优先用 QWEN_BASE_URL(OpenAI 兼容), # DASHSCOPE_BASE_URL 通常是旧版非兼容端点,不作为默认 base_url = ( os.environ.get("QWEN_BASE_URL") or os.environ.get("DASHSCOPE_EMBEDDING_BASE_URL") or DASHSCOPE_DEFAULT_BASE ) m = model or sys_cfg.get("dashscope_model", DASHSCOPE_DEFAULT_MODEL) dimension = int(sys_cfg.get("dimension", DASHSCOPE_DEFAULT_DIM)) batch_size = min(int(sys_cfg.get("batch_size", DASHSCOPE_BATCH_LIMIT)), DASHSCOPE_BATCH_LIMIT) timeout = float(sys_cfg.get("timeout_sec", 30.0)) max_attempts = int(sys_cfg.get("max_attempts", DEFAULT_MAX_ATTEMPTS)) if not api_key: raise EmbeddingError("DASHSCOPE_API_KEY 未配置,请检查 .env") return EmbeddingConfig( provider="dashscope", model=m, api_key=api_key, base_url=base_url, dimension=dimension, batch_size=batch_size, timeout_sec=timeout, max_attempts=max_attempts, ) def make_embedding_client(config: EmbeddingConfig) -> OpenAI: """构造同步 OpenAI 客户端(指向 DashScope 兼容端点)。""" logger.info( "初始化 Embedding 客户端: provider=%s model=%s base_url=%s dim=%d", config.provider, config.model, config.base_url, config.dimension, ) return OpenAI( api_key=config.api_key, base_url=config.base_url, timeout=config.timeout_sec, ) def _chunked(items: list[str], size: int) -> list[list[str]]: """把列表按 size 分块。""" return [items[i : i + size] for i in range(0, len(items), size)] def embed_batch( client: OpenAI, config: EmbeddingConfig, texts: list[str], ) -> list[list[float]]: """批量嵌入,自动分块+重试。 Args: client: OpenAI 客户端 config: Embedding 配置 texts: 待嵌入文本列表 Returns: 与 texts 等长的向量列表,每个为 1024 维 float 列表 """ if not texts: return [] all_results: list[list[float]] = [] chunks = _chunked(texts, config.batch_size) for chunk_idx, chunk in enumerate(chunks): result = _call_with_retry(client, config, chunk, chunk_idx, len(chunks)) all_results.extend(result) return all_results def _call_with_retry( client: OpenAI, config: EmbeddingConfig, batch: list[str], chunk_idx: int, total_chunks: int, ) -> list[list[float]]: """单批嵌入调用,带指数退避重试。""" last_err: Exception | None = None for attempt in range(1, config.max_attempts + 1): try: resp = client.embeddings.create(model=config.model, input=batch) vectors = [d.embedding for d in resp.data] # 维度校验 if vectors and len(vectors[0]) != config.dimension: logger.warning( "实际维度 %d 与预期 %d 不一致", len(vectors[0]), config.dimension, ) logger.debug( "Embedding chunk %d/%d 完成(%d 条,attempt %d)", chunk_idx + 1, total_chunks, len(batch), attempt, ) return vectors except Exception as e: last_err = e logger.warning( "DashScope embed 失败 chunk %d/%d 尝试 %d/%d: %s: %s", chunk_idx + 1, total_chunks, attempt, config.max_attempts, type(e).__name__, e, ) if attempt < config.max_attempts: wait = min(RETRY_BASE_WAIT_SEC * (2 ** (attempt - 1)), RETRY_MAX_WAIT_SEC) time.sleep(wait) raise EmbeddingError( f"DashScope embed 放弃({config.max_attempts} 次): {last_err}", attempts=config.max_attempts, )