codekingpro/portable-devtools
114k
1import json2from typing import Any, AsyncIterator, Dict, Iterator, List, Mapping, Optional3 4import aiohttp5from langchain_core.callbacks import (6 AsyncCallbackManagerForLLMRun,7 CallbackManagerForLLMRun,8)9from langchain_core.language_models.llms import LLM10from langchain_core.outputs import GenerationChunk11from langchain_core.utils import get_from_dict_or_env, pre_init12from pydantic import ConfigDict13 14from langchain_community.utilities.requests import Requests15 16DEFAULT_MODEL_ID = "meta-llama/Meta-Llama-3-70B-Instruct"17 18 19class DeepInfra(LLM):20 """DeepInfra models.21 22 To use, you should have the environment variable ``DEEPINFRA_API_TOKEN``23 set with your API token, or pass it as a named parameter to the24 constructor.25 26 Only supports `text-generation` and `text2text-generation` for now.27 28 Example:29 .. code-block:: python30 31 from langchain_community.llms import DeepInfra32 di = DeepInfra(model_id="google/flan-t5-xl",33 deepinfra_api_token="my-api-key")34 """35 36 model_id: str = DEFAULT_MODEL_ID37 model_kwargs: Optional[Dict] = None38 39 deepinfra_api_token: Optional[str] = None40 41 model_config = ConfigDict(42 extra="forbid",43 )44 45 @pre_init46 def validate_environment(cls, values: Dict) -> Dict:47 """Validate that api key and python package exists in environment."""48 deepinfra_api_token = get_from_dict_or_env(49 values, "deepinfra_api_token", "DEEPINFRA_API_TOKEN"50 )51 values["deepinfra_api_token"] = deepinfra_api_token52 return values53 54 @property55 def _identifying_params(self) -> Mapping[str, Any]:56 """Get the identifying parameters."""57 return {58 **{"model_id": self.model_id},59 **{"model_kwargs": self.model_kwargs},60 }61 62 @property63 def _llm_type(self) -> str:64 """Return type of llm."""65 return "deepinfra"66 67 def _url(self) -> str:68 return f"https://api.deepinfra.com/v1/inference/{self.model_id}"69 70 def _headers(self) -> Dict:71 return {72 "Authorization": f"bearer {self.deepinfra_api_token}",73 "Content-Type": "application/json",74 }75 76 def _body(self, prompt: str, kwargs: Any) -> Dict:77 model_kwargs = self.model_kwargs or {}78 model_kwargs = {**model_kwargs, **kwargs}79 80 return {81 "input": prompt,82 **model_kwargs,83 }84 85 def _handle_status(self, code: int, text: Any) -> None:86 if code >= 500:87 raise Exception(f"DeepInfra Server: Error {text}")88 elif code == 401:89 raise Exception("DeepInfra Server: Unauthorized")90 elif code == 403:91 raise Exception("DeepInfra Server: Unauthorized")92 elif code == 404:93 raise Exception(f"DeepInfra Server: Model not found {self.model_id}")94 elif code == 429:95 raise Exception("DeepInfra Server: Rate limit exceeded")96 elif code >= 400:97 raise ValueError(f"DeepInfra received an invalid payload: {text}")98 elif code != 200:99 raise Exception(100 f"DeepInfra returned an unexpected response with status {code}: {text}"101 )102 103 def _call(104 self,105 prompt: str,106 stop: Optional[List[str]] = None,107 run_manager: Optional[CallbackManagerForLLMRun] = None,108 **kwargs: Any,109 ) -> str:110 """Call out to DeepInfra's inference API endpoint.111 112 Args:113 prompt: The prompt to pass into the model.114 stop: Optional list of stop words to use when generating.115 116 Returns:117 The string generated by the model.118 119 Example:120 .. code-block:: python121 122 response = di("Tell me a joke.")123 """124 125 request = Requests(headers=self._headers())126 response = request.post(url=self._url(), data=self._body(prompt, kwargs))127 128 self._handle_status(response.status_code, response.text)129 data = response.json()130 131 return data["results"][0]["generated_text"]132 133 async def _acall(134 self,135 prompt: str,136 stop: Optional[List[str]] = None,137 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,138 **kwargs: Any,139 ) -> str:140 request = Requests(headers=self._headers())141 async with request.apost(142 url=self._url(), data=self._body(prompt, kwargs)143 ) as response:144 self._handle_status(response.status, response.text)145 data = await response.json()146 return data["results"][0]["generated_text"]147 148 def _stream(149 self,150 prompt: str,151 stop: Optional[List[str]] = None,152 run_manager: Optional[CallbackManagerForLLMRun] = None,153 **kwargs: Any,154 ) -> Iterator[GenerationChunk]:155 request = Requests(headers=self._headers())156 response = request.post(157 url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})158 )159 response_text = response.text160 self._handle_body_errors(response_text)161 self._handle_status(response.status_code, response.text)162 for line in _parse_stream(response.iter_lines()):163 chunk = _handle_sse_line(line)164 if chunk:165 if run_manager:166 run_manager.on_llm_new_token(chunk.text)167 yield chunk168 169 async def _astream(170 self,171 prompt: str,172 stop: Optional[List[str]] = None,173 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,174 **kwargs: Any,175 ) -> AsyncIterator[GenerationChunk]:176 request = Requests(headers=self._headers())177 async with request.apost(178 url=self._url(), data=self._body(prompt, {**kwargs, "stream": True})179 ) as response:180 response_text = await response.text()181 self._handle_body_errors(response_text)182 self._handle_status(response.status, response.text)183 async for line in _parse_stream_async(response.content):184 chunk = _handle_sse_line(line)185 if chunk:186 if run_manager:187 await run_manager.on_llm_new_token(chunk.text)188 yield chunk189 190 def _handle_body_errors(self, body: str) -> None:191 """192 Example error response:193 data: {"error_type": "validation_error",194 "error_message": "ConnectionError: ..."}195 """196 if "error" in body:197 try:198 # Remove data: prefix if present199 if body.startswith("data:"):200 body = body[len("data:") :]201 error_data = json.loads(body)202 error_message = error_data.get("error_message", "Unknown error")203 204 raise Exception(f"DeepInfra Server Error: {error_message}")205 except json.JSONDecodeError:206 raise Exception(f"DeepInfra Server: {body}")207 208 209def _parse_stream(rbody: Iterator[bytes]) -> Iterator[str]:210 for line in rbody:211 _line = _parse_stream_helper(line)212 if _line is not None:213 yield _line214 215 216async def _parse_stream_async(rbody: aiohttp.StreamReader) -> AsyncIterator[str]:217 async for line in rbody:218 _line = _parse_stream_helper(line)219 if _line is not None:220 yield _line221 222 223def _parse_stream_helper(line: bytes) -> Optional[str]:224 if line and line.startswith(b"data:"):225 if line.startswith(b"data: "):226 # SSE event may be valid when it contain whitespace227 line = line[len(b"data: ") :]228 else:229 line = line[len(b"data:") :]230 if line.strip() == b"[DONE]":231 # return here will cause GeneratorExit exception in urllib3232 # and it will close http connection with TCP Reset233 return None234 else:235 return line.decode("utf-8")236 return None237 238 239def _handle_sse_line(line: str) -> Optional[GenerationChunk]:240 try:241 obj = json.loads(line)242 return GenerationChunk(243 text=obj.get("token", {}).get("text"),244 )245 except Exception:246 return None247 