70 lines
2.4 KiB
Python
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]
|