codekingpro/portable-devtools
114k
1import json2import logging3from typing import Any, Dict, Iterator, List, Optional4 5import requests6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models.llms import LLM8from langchain_core.outputs import GenerationChunk9 10logger = logging.getLogger(__name__)11 12 13class CloudflareWorkersAI(LLM):14 """Cloudflare Workers AI service.15 16 To use, you must provide an API token and17 account ID to access Cloudflare Workers AI, and18 pass it as a named parameter to the constructor.19 20 Example:21 .. code-block:: python22 23 from langchain_community.llms.cloudflare_workersai import CloudflareWorkersAI24 25 my_account_id = "my_account_id"26 my_api_token = "my_secret_api_token"27 llm_model = "@cf/meta/llama-2-7b-chat-int8"28 29 cf_ai = CloudflareWorkersAI(30 account_id=my_account_id,31 api_token=my_api_token,32 model=llm_model33 )34 """ # noqa: E50135 36 account_id: str37 api_token: str38 model: str = "@cf/meta/llama-2-7b-chat-int8"39 base_url: str = "https://api.cloudflare.com/client/v4/accounts"40 streaming: bool = False41 endpoint_url: str = ""42 43 def __init__(self, **kwargs: Any) -> None:44 """Initialize the Cloudflare Workers AI class."""45 super().__init__(**kwargs)46 47 self.endpoint_url = f"{self.base_url}/{self.account_id}/ai/run/{self.model}"48 49 @property50 def _llm_type(self) -> str:51 """Return type of LLM."""52 return "cloudflare"53 54 @property55 def _default_params(self) -> Dict[str, Any]:56 """Default parameters"""57 return {}58 59 @property60 def _identifying_params(self) -> Dict[str, Any]:61 """Identifying parameters"""62 return {63 "account_id": self.account_id,64 "api_token": self.api_token,65 "model": self.model,66 "base_url": self.base_url,67 }68 69 def _call_api(self, prompt: str, params: Dict[str, Any]) -> requests.Response:70 """Call Cloudflare Workers API"""71 headers = {"Authorization": f"Bearer {self.api_token}"}72 data = {"prompt": prompt, "stream": self.streaming, **params}73 response = requests.post(74 self.endpoint_url, headers=headers, json=data, stream=self.streaming75 )76 return response77 78 def _process_response(self, response: requests.Response) -> str:79 """Process API response"""80 if response.ok:81 data = response.json()82 return data["result"]["response"]83 else:84 raise ValueError(f"Request failed with status {response.status_code}")85 86 def _stream(87 self,88 prompt: str,89 stop: Optional[List[str]] = None,90 run_manager: Optional[CallbackManagerForLLMRun] = None,91 **kwargs: Any,92 ) -> Iterator[GenerationChunk]:93 """Streaming prediction"""94 original_steaming: bool = self.streaming95 self.streaming = True96 _response_prefix_count = len("data: ")97 _response_stream_end = b"data: [DONE]"98 for chunk in self._call_api(prompt, kwargs).iter_lines():99 if chunk == _response_stream_end:100 break101 if len(chunk) > _response_prefix_count:102 try:103 data = json.loads(chunk[_response_prefix_count:])104 except Exception as e:105 logger.debug(chunk)106 raise e107 if data is not None and "response" in data:108 if run_manager:109 run_manager.on_llm_new_token(data["response"])110 yield GenerationChunk(text=data["response"])111 logger.debug("stream end")112 self.streaming = original_steaming113 114 def _call(115 self,116 prompt: str,117 stop: Optional[List[str]] = None,118 run_manager: Optional[CallbackManagerForLLMRun] = None,119 **kwargs: Any,120 ) -> str:121 """Regular prediction"""122 if self.streaming:123 return "".join(124 [c.text for c in self._stream(prompt, stop, run_manager, **kwargs)]125 )126 else:127 response = self._call_api(prompt, kwargs)128 return self._process_response(response)129 