codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4from typing import Any, Callable, Dict, List, Optional5 6import requests7from langchain_core._api import deprecated8from langchain_core.embeddings import Embeddings9from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init10from pydantic import BaseModel, ConfigDict, SecretStr11from tenacity import (12 before_sleep_log,13 retry,14 stop_after_attempt,15 wait_exponential,16)17 18logger = logging.getLogger(__name__)19 20 21def _create_retry_decorator() -> Callable[[Any], Any]:22 """Returns a tenacity retry decorator."""23 24 multiplier = 125 min_seconds = 126 max_seconds = 427 max_retries = 628 29 return retry(30 reraise=True,31 stop=stop_after_attempt(max_retries),32 wait=wait_exponential(multiplier=multiplier, min=min_seconds, max=max_seconds),33 before_sleep=before_sleep_log(logger, logging.WARNING),34 )35 36 37def embed_with_retry(embeddings: SolarEmbeddings, *args: Any, **kwargs: Any) -> Any:38 """Use tenacity to retry the completion call."""39 retry_decorator = _create_retry_decorator()40 41 @retry_decorator42 def _embed_with_retry(*args: Any, **kwargs: Any) -> Any:43 return embeddings.embed(*args, **kwargs)44 45 return _embed_with_retry(*args, **kwargs)46 47 48@deprecated(49 since="0.0.34", removal="1.0", alternative_import="langchain_upstage.ChatUpstage"50)51class SolarEmbeddings(BaseModel, Embeddings):52 """Solar's embedding service.53 54 To use, you should have the environment variable``SOLAR_API_KEY`` set55 with your API token, or pass it as a named parameter to the constructor.56 57 Example:58 .. code-block:: python59 60 from langchain_community.embeddings import SolarEmbeddings61 embeddings = SolarEmbeddings()62 63 query_text = "This is a test query."64 query_result = embeddings.embed_query(query_text)65 66 document_text = "This is a test document."67 document_result = embeddings.embed_documents([document_text])68 69 """70 71 endpoint_url: str = "https://api.upstage.ai/v1/solar/embeddings"72 """Endpoint URL to use."""73 model: str = "embedding-query"74 """Embeddings model name to use."""75 solar_api_key: Optional[SecretStr] = None76 """API Key for Solar API."""77 78 model_config = ConfigDict(79 extra="forbid",80 )81 82 @pre_init83 def validate_environment(cls, values: Dict) -> Dict:84 """Validate api key exists in environment."""85 solar_api_key = convert_to_secret_str(86 get_from_dict_or_env(values, "solar_api_key", "SOLAR_API_KEY")87 )88 values["solar_api_key"] = solar_api_key89 return values90 91 def embed(92 self,93 text: str,94 ) -> List[List[float]]:95 payload = {96 "model": self.model,97 "input": text,98 }99 100 # HTTP headers for authorization101 headers = {102 "Authorization": f"Bearer {self.solar_api_key.get_secret_value()}", # type: ignore[union-attr]103 "Content-Type": "application/json",104 }105 106 # send request107 response = requests.post(self.endpoint_url, headers=headers, json=payload)108 parsed_response = response.json()109 110 # check for errors111 if len(parsed_response["data"]) == 0:112 raise ValueError(113 f"Solar API returned an error: {parsed_response['base_resp']}"114 )115 116 embedding = parsed_response["data"][0]["embedding"]117 118 return embedding119 120 def embed_documents(self, texts: List[str]) -> List[List[float]]:121 """Embed documents using a Solar embedding endpoint.122 123 Args:124 texts: The list of texts to embed.125 126 Returns:127 List of embeddings, one for each text.128 """129 embeddings = [embed_with_retry(self, text=text) for text in texts]130 return embeddings131 132 def embed_query(self, text: str) -> List[float]:133 """Embed a query using a Solar embedding endpoint.134 135 Args:136 text: The text to embed.137 138 Returns:139 Embeddings for the text.140 """141 embedding = embed_with_retry(self, text=text)142 return embedding143 