codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4from typing import Any, Callable, Dict, List, Optional, Sequence5 6from langchain_core.callbacks import (7 AsyncCallbackManagerForLLMRun,8 CallbackManagerForLLMRun,9)10from langchain_core.language_models.llms import LLM11from langchain_core.load.serializable import Serializable12from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env, pre_init13from pydantic import SecretStr14from tenacity import (15 before_sleep_log,16 retry,17 retry_if_exception_type,18 stop_after_attempt,19 wait_exponential,20)21 22from langchain_community.llms.utils import enforce_stop_tokens23 24logger = logging.getLogger(__name__)25 26 27class _BaseYandexGPT(Serializable):28 iam_token: SecretStr = "" # type: ignore[assignment]29 """Yandex Cloud IAM token for service or user account30 with the `ai.languageModels.user` role"""31 api_key: SecretStr = "" # type: ignore[assignment]32 """Yandex Cloud Api Key for service account33 with the `ai.languageModels.user` role"""34 folder_id: str = ""35 """Yandex Cloud folder ID"""36 model_uri: str = ""37 """Model uri to use."""38 model_name: str = "yandexgpt-lite"39 """Model name to use."""40 model_version: str = "latest"41 """Model version to use."""42 temperature: float = 0.643 """What sampling temperature to use.44 Should be a double number between 0 (inclusive) and 1 (inclusive)."""45 max_tokens: int = 740046 """Sets the maximum limit on the total number of tokens47 used for both the input prompt and the generated response.48 Must be greater than zero and not exceed 7400 tokens."""49 stop: Optional[List[str]] = None50 """Sequences when completion generation will stop."""51 url: str = "llm.api.cloud.yandex.net:443"52 """The url of the API."""53 max_retries: int = 654 """Maximum number of retries to make when generating."""55 sleep_interval: float = 1.056 """Delay between API requests"""57 disable_request_logging: bool = False58 """YandexGPT API logs all request data by default. 59 If you provide personal data, confidential information, disable logging."""60 grpc_metadata: Optional[Sequence] = None61 62 @property63 def _llm_type(self) -> str:64 return "yandex_gpt"65 66 @property67 def _identifying_params(self) -> Dict[str, Any]:68 """Get the identifying parameters."""69 return {70 "model_uri": self.model_uri,71 "temperature": self.temperature,72 "max_tokens": self.max_tokens,73 "stop": self.stop,74 "max_retries": self.max_retries,75 }76 77 @pre_init78 def validate_environment(cls, values: Dict) -> Dict:79 """Validate that iam token exists in environment."""80 81 iam_token = convert_to_secret_str(82 get_from_dict_or_env(values, "iam_token", "YC_IAM_TOKEN", "")83 )84 values["iam_token"] = iam_token85 api_key = convert_to_secret_str(86 get_from_dict_or_env(values, "api_key", "YC_API_KEY", "")87 )88 values["api_key"] = api_key89 folder_id = get_from_dict_or_env(values, "folder_id", "YC_FOLDER_ID", "")90 values["folder_id"] = folder_id91 if api_key.get_secret_value() == "" and iam_token.get_secret_value() == "":92 raise ValueError("Either 'YC_API_KEY' or 'YC_IAM_TOKEN' must be provided.")93 94 if values["iam_token"]:95 values["grpc_metadata"] = [96 ("authorization", f"Bearer {values['iam_token'].get_secret_value()}")97 ]98 if values["folder_id"]:99 values["grpc_metadata"].append(("x-folder-id", values["folder_id"]))100 else:101 values["grpc_metadata"] = [102 ("authorization", f"Api-Key {values['api_key'].get_secret_value()}"),103 ]104 if values["model_uri"] == "" and values["folder_id"] == "":105 raise ValueError("Either 'model_uri' or 'folder_id' must be provided.")106 if not values["model_uri"]:107 values["model_uri"] = (108 f"gpt://{values['folder_id']}/{values['model_name']}/{values['model_version']}"109 )110 if values["disable_request_logging"]:111 values["grpc_metadata"].append(112 (113 "x-data-logging-enabled",114 "false",115 )116 )117 return values118 119 120class YandexGPT(_BaseYandexGPT, LLM):121 """Yandex large language models.122 123 To use, you should have the ``yandexcloud`` python package installed.124 125 There are two authentication options for the service account126 with the ``ai.languageModels.user`` role:127 - You can specify the token in a constructor parameter `iam_token`128 or in an environment variable `YC_IAM_TOKEN`.129 - You can specify the key in a constructor parameter `api_key`130 or in an environment variable `YC_API_KEY`.131 132 To use the default model specify the folder ID in a parameter `folder_id`133 or in an environment variable `YC_FOLDER_ID`.134 135 Or specify the model URI in a constructor parameter `model_uri`136 137 Example:138 .. code-block:: python139 140 from langchain_community.llms import YandexGPT141 yandex_gpt = YandexGPT(iam_token="t1.9eu...", folder_id="b1g...")142 """143 144 def _call(145 self,146 prompt: str,147 stop: Optional[List[str]] = None,148 run_manager: Optional[CallbackManagerForLLMRun] = None,149 **kwargs: Any,150 ) -> str:151 """Call the Yandex GPT model and return the output.152 153 Args:154 prompt: The prompt to pass into the model.155 stop: Optional list of stop words to use when generating.156 157 Returns:158 The string generated by the model.159 160 Example:161 .. code-block:: python162 163 response = YandexGPT("Tell me a joke.")164 """165 text = completion_with_retry(self, prompt=prompt)166 if stop is not None:167 text = enforce_stop_tokens(text, stop)168 return text169 170 async def _acall(171 self,172 prompt: str,173 stop: Optional[List[str]] = None,174 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,175 **kwargs: Any,176 ) -> str:177 """Async call the Yandex GPT model and return the output.178 179 Args:180 prompt: The prompt to pass into the model.181 stop: Optional list of stop words to use when generating.182 183 Returns:184 The string generated by the model.185 """186 text = await acompletion_with_retry(self, prompt=prompt)187 if stop is not None:188 text = enforce_stop_tokens(text, stop)189 return text190 191 192def _make_request(193 self: YandexGPT,194 prompt: str,195) -> str:196 try:197 import grpc198 from google.protobuf.wrappers_pb2 import DoubleValue, Int64Value199 200 try:201 from yandex.cloud.ai.foundation_models.v1.text_common_pb2 import (202 CompletionOptions,203 Message,204 )205 from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2 import ( # noqa: E501206 CompletionRequest,207 )208 from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2_grpc import ( # noqa: E501209 TextGenerationServiceStub,210 )211 except ModuleNotFoundError:212 from yandex.cloud.ai.foundation_models.v1.foundation_models_pb2 import (213 CompletionOptions,214 Message,215 )216 from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2 import ( # noqa: E501217 CompletionRequest,218 )219 from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2_grpc import ( # noqa: E501220 TextGenerationServiceStub,221 )222 except ImportError as e:223 raise ImportError(224 "Please install YandexCloud SDK with `pip install yandexcloud` \225 or upgrade it to recent version."226 ) from e227 channel_credentials = grpc.ssl_channel_credentials()228 channel = grpc.secure_channel(self.url, channel_credentials)229 request = CompletionRequest(230 model_uri=self.model_uri,231 completion_options=CompletionOptions(232 temperature=DoubleValue(value=self.temperature),233 max_tokens=Int64Value(value=self.max_tokens),234 ),235 messages=[Message(role="user", text=prompt)],236 )237 stub = TextGenerationServiceStub(channel)238 res = stub.Completion(request, metadata=self.grpc_metadata)239 return list(res)[0].alternatives[0].message.text240 241 242async def _amake_request(self: YandexGPT, prompt: str) -> str:243 try:244 import asyncio245 246 import grpc247 from google.protobuf.wrappers_pb2 import DoubleValue, Int64Value248 249 try:250 from yandex.cloud.ai.foundation_models.v1.text_common_pb2 import (251 CompletionOptions,252 Message,253 )254 from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2 import ( # noqa: E501255 CompletionRequest,256 CompletionResponse,257 )258 from yandex.cloud.ai.foundation_models.v1.text_generation.text_generation_service_pb2_grpc import ( # noqa: E501259 TextGenerationAsyncServiceStub,260 )261 except ModuleNotFoundError:262 from yandex.cloud.ai.foundation_models.v1.foundation_models_pb2 import (263 CompletionOptions,264 Message,265 )266 from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2 import ( # noqa: E501267 CompletionRequest,268 CompletionResponse,269 )270 from yandex.cloud.ai.foundation_models.v1.foundation_models_service_pb2_grpc import ( # noqa: E501271 TextGenerationAsyncServiceStub,272 )273 from yandex.cloud.operation.operation_service_pb2 import GetOperationRequest274 from yandex.cloud.operation.operation_service_pb2_grpc import (275 OperationServiceStub,276 )277 except ImportError as e:278 raise ImportError(279 "Please install YandexCloud SDK with `pip install yandexcloud` \280 or upgrade it to recent version."281 ) from e282 operation_api_url = "operation.api.cloud.yandex.net:443"283 channel_credentials = grpc.ssl_channel_credentials()284 async with grpc.aio.secure_channel(self.url, channel_credentials) as channel:285 request = CompletionRequest(286 model_uri=self.model_uri,287 completion_options=CompletionOptions(288 temperature=DoubleValue(value=self.temperature),289 max_tokens=Int64Value(value=self.max_tokens),290 ),291 messages=[Message(role="user", text=prompt)],292 )293 stub = TextGenerationAsyncServiceStub(channel)294 operation = await stub.Completion(request, metadata=self.grpc_metadata)295 async with grpc.aio.secure_channel(296 operation_api_url, channel_credentials297 ) as operation_channel:298 operation_stub = OperationServiceStub(operation_channel)299 while not operation.done:300 await asyncio.sleep(1)301 operation_request = GetOperationRequest(operation_id=operation.id)302 operation = await operation_stub.Get(303 operation_request,304 metadata=self.grpc_metadata,305 )306 307 completion_response = CompletionResponse()308 operation.response.Unpack(completion_response)309 return completion_response.alternatives[0].message.text310 311 312def _create_retry_decorator(llm: YandexGPT) -> Callable[[Any], Any]:313 from grpc import RpcError314 315 min_seconds = llm.sleep_interval316 max_seconds = 60317 return retry(318 reraise=True,319 stop=stop_after_attempt(llm.max_retries),320 wait=wait_exponential(multiplier=1, min=min_seconds, max=max_seconds),321 retry=(retry_if_exception_type((RpcError))),322 before_sleep=before_sleep_log(logger, logging.WARNING),323 )324 325 326def completion_with_retry(llm: YandexGPT, **kwargs: Any) -> Any:327 """Use tenacity to retry the completion call."""328 retry_decorator = _create_retry_decorator(llm)329 330 @retry_decorator331 def _completion_with_retry(**_kwargs: Any) -> Any:332 return _make_request(llm, **_kwargs)333 334 return _completion_with_retry(**kwargs)335 336 337async def acompletion_with_retry(llm: YandexGPT, **kwargs: Any) -> Any:338 """Use tenacity to retry the async completion call."""339 retry_decorator = _create_retry_decorator(llm)340 341 @retry_decorator342 async def _completion_with_retry(**_kwargs: Any) -> Any:343 return await _amake_request(llm, **_kwargs)344 345 return await _completion_with_retry(**kwargs)346 