from typing import List import time import logging from langchain_core.embeddings import Embeddings from openai import OpenAI, RateLimitError from app.config import settings logger = logging.getLogger(__name__) class SiliconFlowEmbeddings(Embeddings): """OpenAI 兼容 embeddings 直连客户端(绕开 langchain tiktoken token 化) langchain 的 OpenAIEmbeddings 默认 tiktoken_enabled=True,会把文本编码成 整数 token 数组发送给 API(OpenAI 官方接受,但 siliconflow 等第三方网关 只接受字符串 input,导致 400 code:20015 parameter invalid)。 本类直接用 openai SDK 发送字符串数组,兼容 siliconflow 等网关。 内置 429 TPM 指数退避重试:SiliconFlow 全量同步易触 TPM 限额,429 是 临时性状态,退避重试能自动恢复完成,避免 sync 整段中断。 """ _BATCH_SIZE = 32 _THROTTLE_SECONDS = 0.02 _MAX_RETRY = 3 _INITIAL_RETRY_SECONDS = 2.0 def __init__(self, model: str, api_key: str, base_url: str): self.model = model self.client = OpenAI(api_key=api_key, base_url=base_url) def _call_with_retry(self, batch): wait = self._INITIAL_RETRY_SECONDS for attempt in range(self._MAX_RETRY + 1): try: return self.client.embeddings.create(model=self.model, input=batch) except RateLimitError as e: if attempt == self._MAX_RETRY: raise logger.warning("embeddings 429 TPM 限制,%.0fs 后重试 (%d/%d)", wait, attempt + 1, self._MAX_RETRY) time.sleep(wait) wait = min(wait * 2, 30.0) raise RuntimeError("unreachable") def embed_documents(self, texts: List[str]) -> List[List[float]]: result = [] for i in range(0, len(texts), self._BATCH_SIZE): batch = texts[i:i + self._BATCH_SIZE] resp = self._call_with_retry(batch) result.extend(d.embedding for d in resp.data) if i + self._BATCH_SIZE < len(texts): time.sleep(self._THROTTLE_SECONDS) return result def embed_query(self, text: str) -> List[float]: resp = self._call_with_retry([text]) return resp.data[0].embedding _embeddings = None def get_embeddings(): global _embeddings if _embeddings is None: _embeddings = SiliconFlowEmbeddings( model=settings.embedding_model, api_key=settings.effective_embedding_api_key, base_url=settings.effective_embedding_base_url, ) return _embeddings