Files
kbqa-system/backend/src/kbqa/rag/embeddings.py
T

70 lines
2.4 KiB
Python

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]