Underground-Digital/Workflow-Engine
0
1import time2from json import JSONDecodeError, dumps3from typing import Optional4 5import requests6 7from core.entities.embedding_type import EmbeddingInputType8from core.model_runtime.entities.common_entities import I18nObject9from core.model_runtime.entities.model_entities import AIModelEntity, FetchFrom, ModelPropertyKey, ModelType, PriceType10from core.model_runtime.entities.text_embedding_entities import EmbeddingUsage, TextEmbeddingResult11from core.model_runtime.errors.invoke import (12 InvokeAuthorizationError,13 InvokeBadRequestError,14 InvokeConnectionError,15 InvokeError,16 InvokeRateLimitError,17 InvokeServerUnavailableError,18)19from core.model_runtime.errors.validate import CredentialsValidateFailedError20from core.model_runtime.model_providers.__base.text_embedding_model import TextEmbeddingModel21 22 23class MixedBreadTextEmbeddingModel(TextEmbeddingModel):24 """25 Model class for MixedBread text embedding model.26 """27 28 api_base: str = "https://api.mixedbread.ai/v1"29 30 def _invoke(31 self,32 model: str,33 credentials: dict,34 texts: list[str],35 user: Optional[str] = None,36 input_type: EmbeddingInputType = EmbeddingInputType.DOCUMENT,37 ) -> TextEmbeddingResult:38 """39 Invoke text embedding model40 41 :param model: model name42 :param credentials: model credentials43 :param texts: texts to embed44 :param user: unique user id45 :param input_type: input type46 :return: embeddings result47 """48 api_key = credentials["api_key"]49 if not api_key:50 raise CredentialsValidateFailedError("api_key is required")51 52 base_url = credentials.get("base_url", self.api_base)53 base_url = base_url.removesuffix("/")54 55 url = base_url + "/embeddings"56 headers = {"Authorization": "Bearer " + api_key, "Content-Type": "application/json"}57 58 data = {"model": model, "input": texts}59 60 try:61 response = requests.post(url, headers=headers, data=dumps(data))62 except Exception as e:63 raise InvokeConnectionError(str(e))64 65 if response.status_code != 200:66 try:67 resp = response.json()68 msg = resp["detail"]69 if response.status_code == 401:70 raise InvokeAuthorizationError(msg)71 elif response.status_code == 429:72 raise InvokeRateLimitError(msg)73 elif response.status_code == 500:74 raise InvokeServerUnavailableError(msg)75 else:76 raise InvokeBadRequestError(msg)77 except JSONDecodeError as e:78 raise InvokeServerUnavailableError(79 f"Failed to convert response to json: {e} with text: {response.text}"80 )81 82 try:83 resp = response.json()84 embeddings = resp["data"]85 usage = resp["usage"]86 except Exception as e:87 raise InvokeServerUnavailableError(f"Failed to convert response to json: {e} with text: {response.text}")88 89 usage = self._calc_response_usage(model=model, credentials=credentials, tokens=usage["total_tokens"])90 91 result = TextEmbeddingResult(92 model=model, embeddings=[[float(data) for data in x["embedding"]] for x in embeddings], usage=usage93 )94 95 return result96 97 def get_num_tokens(self, model: str, credentials: dict, texts: list[str]) -> int:98 """99 Get number of tokens for given prompt messages100 101 :param model: model name102 :param credentials: model credentials103 :param texts: texts to embed104 :return:105 """106 return sum(self._get_num_tokens_by_gpt2(text) for text in texts)107 108 def validate_credentials(self, model: str, credentials: dict) -> None:109 """110 Validate model credentials111 112 :param model: model name113 :param credentials: model credentials114 :return:115 """116 try:117 self._invoke(model=model, credentials=credentials, texts=["ping"])118 except Exception as e:119 raise CredentialsValidateFailedError(f"Credentials validation failed: {e}")120 121 @property122 def _invoke_error_mapping(self) -> dict[type[InvokeError], list[type[Exception]]]:123 return {124 InvokeConnectionError: [InvokeConnectionError],125 InvokeServerUnavailableError: [InvokeServerUnavailableError],126 InvokeRateLimitError: [InvokeRateLimitError],127 InvokeAuthorizationError: [InvokeAuthorizationError],128 InvokeBadRequestError: [KeyError, InvokeBadRequestError],129 }130 131 def _calc_response_usage(self, model: str, credentials: dict, tokens: int) -> EmbeddingUsage:132 """133 Calculate response usage134 135 :param model: model name136 :param credentials: model credentials137 :param tokens: input tokens138 :return: usage139 """140 # get input price info141 input_price_info = self.get_price(142 model=model, credentials=credentials, price_type=PriceType.INPUT, tokens=tokens143 )144 145 # transform usage146 usage = EmbeddingUsage(147 tokens=tokens,148 total_tokens=tokens,149 unit_price=input_price_info.unit_price,150 price_unit=input_price_info.unit,151 total_price=input_price_info.total_amount,152 currency=input_price_info.currency,153 latency=time.perf_counter() - self.started_at,154 )155 156 return usage157 158 def get_customizable_model_schema(self, model: str, credentials: dict) -> AIModelEntity:159 """160 generate custom model entities from credentials161 """162 entity = AIModelEntity(163 model=model,164 label=I18nObject(en_US=model),165 model_type=ModelType.TEXT_EMBEDDING,166 fetch_from=FetchFrom.CUSTOMIZABLE_MODEL,167 model_properties={ModelPropertyKey.CONTEXT_SIZE: int(credentials.get("context_size", "512"))},168 )169 170 return entity171 