codekingpro/portable-devtools
114k
1import json2import logging3import time4from typing import Any, List5 6import requests7from langchain_core.embeddings import Embeddings8from pydantic import BaseModel, ConfigDict9 10logger = logging.getLogger(__name__)11 12 13class OVHCloudEmbeddings(BaseModel, Embeddings):14 """15 OVHcloud AI Endpoints Embeddings.16 """17 18 """ OVHcloud AI Endpoints Access Token"""19 access_token: str = ""20 21 """ OVHcloud AI Endpoints model name for embeddings generation"""22 model_name: str = ""23 24 """ OVHcloud AI Endpoints region"""25 region: str = "kepler"26 27 model_config = ConfigDict(extra="forbid", protected_namespaces=())28 29 def __init__(self, **kwargs: Any):30 super().__init__(**kwargs)31 if self.access_token == "":32 raise ValueError("Access token is required for OVHCloud embeddings.")33 if self.model_name == "":34 raise ValueError("Model name is required for OVHCloud embeddings.")35 if self.region == "":36 raise ValueError("Region is required for OVHCloud embeddings.")37 38 def _generate_embedding(self, text: str) -> List[float]:39 """Generate embeddings from OVHCLOUD AIE.40 Args:41 text (str): The text to embed.42 Returns:43 List[float]: Embeddings for the text.44 """45 46 return self._send_request_to_ai_endpoints("text/plain", text, "text2vec")47 48 def embed_documents(self, texts: List[str]) -> List[List[float]]:49 """Embed a list of documents.50 Args:51 texts (List[str]): The list of texts to embed.52 53 Returns:54 List[List[float]]: List of embeddings, one for each input text.55 56 """57 58 return self._send_request_to_ai_endpoints(59 "application/json", json.dumps(texts), "batch_text2vec"60 )61 62 def embed_query(self, text: str) -> List[float]:63 """Embed a single query text.64 Args:65 text (str): The text to embed.66 Returns:67 List[float]: Embeddings for the text.68 """69 return self._generate_embedding(text)70 71 def _send_request_to_ai_endpoints(72 self, contentType: str, payload: str, route: str73 ) -> Any:74 """Send a HTTPS request to OVHcloud AI Endpoints75 Args:76 contentType (str): The content type of the request, application/json or text/plain.77 payload (str): The payload of the request.78 route (str): The route of the request, batch_text2vec or text2vec.79 """ # noqa: E50180 headers = {81 "content-type": contentType,82 "Authorization": f"Bearer {self.access_token}",83 }84 85 session = requests.session()86 while True:87 response = session.post(88 (89 f"https://{self.model_name}.endpoints.{self.region}"90 f".ai.cloud.ovh.net/api/{route}"91 ),92 headers=headers,93 data=payload,94 )95 if response.status_code != 200:96 if response.status_code == 429:97 """Rate limit exceeded, wait for reset"""98 reset_time = int(response.headers.get("RateLimit-Reset", 0))99 logger.info("Rate limit exceeded. Waiting %d seconds.", reset_time)100 if reset_time > 0:101 time.sleep(reset_time)102 continue103 else:104 """Rate limit reset time has passed, retry immediately"""105 continue106 if response.status_code == 401:107 """ Unauthorized, retry with new token """108 raise ValueError("Unauthorized, retry with new token")109 """ Handle other non-200 status codes """110 raise ValueError(111 "Request failed with status code: {status_code}, {text}".format(112 status_code=response.status_code, text=response.text113 )114 )115 return response.json()116 