codekingpro/portable-devtools
114k
1"""Wrapper around EdenAI's Generation API."""2 3import logging4from typing import Any, Dict, List, Literal, Optional5 6from aiohttp import ClientSession7from langchain_core.callbacks import (8 AsyncCallbackManagerForLLMRun,9 CallbackManagerForLLMRun,10)11from langchain_core.language_models.llms import LLM12from langchain_core.utils import get_from_dict_or_env, pre_init13from langchain_core.utils.pydantic import get_fields14from pydantic import ConfigDict, Field, model_validator15 16from langchain_community.llms.utils import enforce_stop_tokens17from langchain_community.utilities.requests import Requests18 19logger = logging.getLogger(__name__)20 21 22class EdenAI(LLM):23 """EdenAI models.24 25 To use, you should have26 the environment variable ``EDENAI_API_KEY`` set with your API token.27 You can find your token here: https://app.edenai.run/admin/account/settings28 29 `feature` and `subfeature` are required, but any other model parameters can also be30 passed in with the format params={model_param: value, ...}31 32 for api reference check edenai documentation: http://docs.edenai.co.33 """34 35 base_url: str = "https://api.edenai.run/v2"36 37 edenai_api_key: Optional[str] = None38 39 feature: Literal["text", "image"] = "text"40 """Which generative feature to use, use text by default"""41 42 subfeature: Literal["generation"] = "generation"43 """Subfeature of above feature, use generation by default"""44 45 provider: str46 """Generative provider to use (eg: openai,stabilityai,cohere,google etc.)"""47 48 model: Optional[str] = None49 """50 model name for above provider (eg: 'gpt-3.5-turbo-instruct' for openai)51 available models are shown on https://docs.edenai.co/ under 'available providers'52 """53 54 # Optional parameters to add depending of chosen feature55 # see api reference for more infos56 temperature: Optional[float] = Field(default=None, ge=0, le=1) # for text57 max_tokens: Optional[int] = Field(default=None, ge=0) # for text58 resolution: Optional[Literal["256x256", "512x512", "1024x1024"]] = None # for image59 60 params: Dict[str, Any] = Field(default_factory=dict)61 """62 DEPRECATED: use temperature, max_tokens, resolution directly63 optional parameters to pass to api64 """65 66 model_kwargs: Dict[str, Any] = Field(default_factory=dict)67 """extra parameters"""68 69 stop_sequences: Optional[List[str]] = None70 """Stop sequences to use."""71 72 model_config = ConfigDict(73 extra="forbid",74 )75 76 @pre_init77 def validate_environment(cls, values: Dict) -> Dict:78 """Validate that api key exists in environment."""79 values["edenai_api_key"] = get_from_dict_or_env(80 values, "edenai_api_key", "EDENAI_API_KEY"81 )82 return values83 84 @model_validator(mode="before")85 @classmethod86 def build_extra(cls, values: Dict[str, Any]) -> Any:87 """Build extra kwargs from additional params that were passed in."""88 all_required_field_names = {field.alias for field in get_fields(cls).values()}89 90 extra = values.get("model_kwargs", {})91 for field_name in list(values):92 if field_name not in all_required_field_names:93 if field_name in extra:94 raise ValueError(f"Found {field_name} supplied twice.")95 logger.warning(96 f"""{field_name} was transferred to model_kwargs.97 Please confirm that {field_name} is what you intended."""98 )99 extra[field_name] = values.pop(field_name)100 values["model_kwargs"] = extra101 return values102 103 @property104 def _llm_type(self) -> str:105 """Return type of model."""106 return "edenai"107 108 def _format_output(self, output: dict) -> str:109 if self.feature == "text":110 return output[self.provider]["generated_text"]111 else:112 return output[self.provider]["items"][0]["image"]113 114 @staticmethod115 def get_user_agent() -> str:116 from langchain_community import __version__117 118 return f"langchain/{__version__}"119 120 def _call(121 self,122 prompt: str,123 stop: Optional[List[str]] = None,124 run_manager: Optional[CallbackManagerForLLMRun] = None,125 **kwargs: Any,126 ) -> str:127 """Call out to EdenAI's text generation endpoint.128 129 Args:130 prompt: The prompt to pass into the model.131 132 Returns:133 json formatted str response.134 """135 stops = None136 if self.stop_sequences is not None and stop is not None:137 raise ValueError(138 "stop sequences found in both the input and default params."139 )140 elif self.stop_sequences is not None:141 stops = self.stop_sequences142 else:143 stops = stop144 145 url = f"{self.base_url}/{self.feature}/{self.subfeature}"146 headers = {147 "Authorization": f"Bearer {self.edenai_api_key}",148 "User-Agent": self.get_user_agent(),149 }150 payload: Dict[str, Any] = {151 "providers": self.provider,152 "text": prompt,153 "max_tokens": self.max_tokens,154 "temperature": self.temperature,155 "resolution": self.resolution,156 **self.params,157 **kwargs,158 "num_images": 1, # always limit to 1 (ignored for text)159 }160 161 # filter None values to not pass them to the http payload162 payload = {k: v for k, v in payload.items() if v is not None}163 164 if self.model is not None:165 payload["settings"] = {self.provider: self.model}166 167 request = Requests(headers=headers)168 response = request.post(url=url, data=payload)169 170 if response.status_code >= 500:171 raise Exception(f"EdenAI Server: Error {response.status_code}")172 elif response.status_code >= 400:173 raise ValueError(f"EdenAI received an invalid payload: {response.text}")174 elif response.status_code != 200:175 raise Exception(176 f"EdenAI returned an unexpected response with status "177 f"{response.status_code}: {response.text}"178 )179 180 data = response.json()181 provider_response = data[self.provider]182 if provider_response.get("status") == "fail":183 err_msg = provider_response.get("error", {}).get("message")184 raise Exception(err_msg)185 186 output = self._format_output(data)187 188 if stops is not None:189 output = enforce_stop_tokens(output, stops)190 191 return output192 193 async def _acall(194 self,195 prompt: str,196 stop: Optional[List[str]] = None,197 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,198 **kwargs: Any,199 ) -> str:200 """Call EdenAi model to get predictions based on the prompt.201 202 Args:203 prompt: The prompt to pass into the model.204 stop: A list of stop words (optional).205 run_manager: A callback manager for async interaction with LLMs.206 207 Returns:208 The string generated by the model.209 """210 211 stops = None212 if self.stop_sequences is not None and stop is not None:213 raise ValueError(214 "stop sequences found in both the input and default params."215 )216 elif self.stop_sequences is not None:217 stops = self.stop_sequences218 else:219 stops = stop220 221 url = f"{self.base_url}/{self.feature}/{self.subfeature}"222 headers = {223 "Authorization": f"Bearer {self.edenai_api_key}",224 "User-Agent": self.get_user_agent(),225 }226 payload: Dict[str, Any] = {227 "providers": self.provider,228 "text": prompt,229 "max_tokens": self.max_tokens,230 "temperature": self.temperature,231 "resolution": self.resolution,232 **self.params,233 **kwargs,234 "num_images": 1, # always limit to 1 (ignored for text)235 }236 237 # filter `None` values to not pass them to the http payload as null238 payload = {k: v for k, v in payload.items() if v is not None}239 240 if self.model is not None:241 payload["settings"] = {self.provider: self.model}242 243 async with ClientSession() as session:244 async with session.post(url, json=payload, headers=headers) as response:245 if response.status >= 500:246 raise Exception(f"EdenAI Server: Error {response.status}")247 elif response.status >= 400:248 raise ValueError(249 f"EdenAI received an invalid payload: {response.text}"250 )251 elif response.status != 200:252 raise Exception(253 f"EdenAI returned an unexpected response with status "254 f"{response.status}: {response.text}"255 )256 257 response_json = await response.json()258 provider_response = response_json[self.provider]259 if provider_response.get("status") == "fail":260 err_msg = provider_response.get("error", {}).get("message")261 raise Exception(err_msg)262 263 output = self._format_output(response_json)264 if stops is not None:265 output = enforce_stop_tokens(output, stops)266 267 return output268 