import logging from langchain_openai import OpenAIEmbeddings from kbqa.api.errors import AppError from kbqa.config import Settings, validate_live_settings logger = logging.getLogger(__name__) class EmbeddingProvider: def __init__(self, settings: Settings) -> None: self.settings = settings self._client: OpenAIEmbeddings | None = None def _get_client(self) -> OpenAIEmbeddings: if self._client is None: validate_live_settings(self.settings) self._client = OpenAIEmbeddings( api_key=self.settings.dashscope_api_key, base_url=self.settings.dashscope_base_url, model=self.settings.embedding_model, dimensions=self.settings.embedding_dim, chunk_size=10, max_retries=2, timeout=60.0, check_embedding_ctx_length=False, ) return self._client def _validate_vectors(self, vectors: list[list[float]]) -> list[list[float]]: if any(len(vector) != self.settings.embedding_dim for vector in vectors): raise AppError( "EMBEDDING_DIMENSION_MISMATCH", "Embedding 返回的向量维度不正确", 502, retryable=True, ) return vectors async def embed_documents(self, texts: list[str]) -> list[list[float]]: try: vectors = await self._get_client().aembed_documents(texts) except AppError: raise except Exception as exc: logger.exception("Embedding document request failed error_type=%s", type(exc).__name__) raise AppError( "EMBEDDING_UPSTREAM_ERROR", "文档向量化失败,请稍后重试", 502, retryable=True, ) from exc return self._validate_vectors(vectors) async def embed_query(self, text: str) -> list[float]: try: vector = await self._get_client().aembed_query(text) except AppError: raise except Exception as exc: logger.exception("Embedding query request failed error_type=%s", type(exc).__name__) raise AppError( "EMBEDDING_UPSTREAM_ERROR", "查询向量化失败,请稍后重试", 502, retryable=True, ) from exc return self._validate_vectors([vector])[0]