codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Mapping, Optional2 3import requests4from langchain_core.embeddings import Embeddings5from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init6from pydantic import BaseModel, ConfigDict, SecretStr7from requests.adapters import HTTPAdapter, Retry8from typing_extensions import NotRequired, TypedDict9 10# Currently supported maximum batch size for embedding requests11MAX_BATCH_SIZE = 25612EMBAAS_API_URL = "https://api.embaas.io/v1/embeddings/"13 14 15class EmbaasEmbeddingsPayload(TypedDict):16 """Payload for the Embaas embeddings API."""17 18 model: str19 texts: List[str]20 instruction: NotRequired[str]21 22 23class EmbaasEmbeddings(BaseModel, Embeddings):24 """Embaas's embedding service.25 26 To use, you should have the27 environment variable ``EMBAAS_API_KEY`` set with your API key, or pass28 it as a named parameter to the constructor.29 30 Example:31 .. code-block:: python32 33 # initialize with default model and instruction34 from langchain_community.embeddings import EmbaasEmbeddings35 emb = EmbaasEmbeddings()36 37 # initialize with custom model and instruction38 from langchain_community.embeddings import EmbaasEmbeddings39 emb_model = "instructor-large"40 emb_inst = "Represent the Wikipedia document for retrieval"41 emb = EmbaasEmbeddings(42 model=emb_model,43 instruction=emb_inst44 )45 """46 47 model: str = "e5-large-v2"48 """The model used for embeddings."""49 instruction: Optional[str] = None50 """Instruction used for domain-specific embeddings."""51 api_url: str = EMBAAS_API_URL52 """The URL for the embaas embeddings API."""53 embaas_api_key: Optional[SecretStr] = None54 """max number of retries for requests"""55 max_retries: Optional[int] = 356 """request timeout in seconds"""57 timeout: Optional[int] = 3058 59 model_config = ConfigDict(60 extra="forbid",61 )62 63 @pre_init64 def validate_environment(cls, values: Dict) -> Dict:65 """Validate that api key and python package exists in environment."""66 embaas_api_key = convert_to_secret_str(67 get_from_dict_or_env(values, "embaas_api_key", "EMBAAS_API_KEY")68 )69 values["embaas_api_key"] = embaas_api_key70 return values71 72 @property73 def _identifying_params(self) -> Mapping[str, Any]:74 """Get the identifying params."""75 return {"model": self.model, "instruction": self.instruction}76 77 def _generate_payload(self, texts: List[str]) -> EmbaasEmbeddingsPayload:78 """Generates payload for the API request."""79 payload = EmbaasEmbeddingsPayload(texts=texts, model=self.model)80 if self.instruction:81 payload["instruction"] = self.instruction82 return payload83 84 def _handle_request(self, payload: EmbaasEmbeddingsPayload) -> List[List[float]]:85 """Sends a request to the Embaas API and handles the response."""86 headers = {87 "Authorization": f"Bearer {self.embaas_api_key.get_secret_value()}", # type: ignore[union-attr]88 "Content-Type": "application/json",89 }90 91 session = requests.Session()92 retries = Retry(93 total=self.max_retries,94 backoff_factor=0.5,95 allowed_methods=["POST"],96 raise_on_status=True,97 )98 99 session.mount("http://", HTTPAdapter(max_retries=retries))100 session.mount("https://", HTTPAdapter(max_retries=retries))101 response = session.post(102 self.api_url,103 headers=headers,104 json=payload,105 timeout=self.timeout,106 )107 108 parsed_response = response.json()109 embeddings = [item["embedding"] for item in parsed_response["data"]]110 111 return embeddings112 113 def _generate_embeddings(self, texts: List[str]) -> List[List[float]]:114 """Generate embeddings using the Embaas API."""115 payload = self._generate_payload(texts)116 try:117 return self._handle_request(payload)118 except requests.exceptions.RequestException as e:119 if e.response is None or not e.response.text:120 raise ValueError(f"Error raised by embaas embeddings API: {e}")121 122 parsed_response = e.response.json()123 if "message" in parsed_response:124 raise ValueError(125 "Validation Error raised by embaas embeddings API:"126 f"{parsed_response['message']}"127 )128 raise129 130 def embed_documents(self, texts: List[str]) -> List[List[float]]:131 """Get embeddings for a list of texts.132 133 Args:134 texts: The list of texts to get embeddings for.135 136 Returns:137 List of embeddings, one for each text.138 """139 batches = [140 texts[i : i + MAX_BATCH_SIZE] for i in range(0, len(texts), MAX_BATCH_SIZE)141 ]142 embeddings = [self._generate_embeddings(batch) for batch in batches]143 # flatten the list of lists into a single list144 return [embedding for batch in embeddings for embedding in batch]145 146 def embed_query(self, text: str) -> List[float]:147 """Get embeddings for a single text.148 149 Args:150 text: The text to get embeddings for.151 152 Returns:153 List of embeddings.154 """155 return self.embed_documents([text])[0]156 