codekingpro/portable-devtools
114k
1from __future__ import annotations2 3import logging4from typing import (5 TYPE_CHECKING,6 Any,7 AsyncGenerator,8 AsyncIterator,9 Callable,10 Dict,11 Generator,12 Iterator,13 List,14 Mapping,15 Optional,16 Tuple,17 Type,18 Union,19)20 21from langchain_core.callbacks import (22 AsyncCallbackManagerForLLMRun,23 CallbackManagerForLLMRun,24)25from langchain_core.language_models.chat_models import (26 BaseChatModel,27 agenerate_from_stream,28 generate_from_stream,29)30from langchain_core.language_models.llms import create_base_retry_decorator31from langchain_core.messages import AIMessageChunk, BaseMessage, BaseMessageChunk32from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult33from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env34from pydantic import BaseModel, Field, SecretStr, model_validator35from typing_extensions import Self36 37from langchain_community.adapters.openai import (38 convert_dict_to_message,39 convert_message_to_dict,40)41from langchain_community.chat_models.openai import _convert_delta_to_message_chunk42 43if TYPE_CHECKING:44 from gpt_router.models import ChunkedGenerationResponse, GenerationResponse45 46 47logger = logging.getLogger(__name__)48 49DEFAULT_API_BASE_URL = "https://gpt-router-preview.writesonic.com"50 51 52class GPTRouterException(Exception):53 """Error with the `GPTRouter APIs`"""54 55 56class GPTRouterModel(BaseModel):57 """GPTRouter model."""58 59 name: str60 provider_name: str61 62 63def get_ordered_generation_requests(64 models_priority_list: List[GPTRouterModel], **kwargs: Any65) -> List:66 """67 Return the body for the model router input.68 """69 70 from gpt_router.models import GenerationParams, ModelGenerationRequest71 72 return [73 ModelGenerationRequest(74 model_name=model.name,75 provider_name=model.provider_name,76 order=index + 1,77 prompt_params=GenerationParams(**kwargs),78 )79 for index, model in enumerate(models_priority_list)80 ]81 82 83def _create_retry_decorator(84 llm: GPTRouter,85 run_manager: Optional[86 Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]87 ] = None,88) -> Callable[[Any], Any]:89 from gpt_router import exceptions90 91 errors = [92 exceptions.GPTRouterApiTimeoutError,93 exceptions.GPTRouterInternalServerError,94 exceptions.GPTRouterNotAvailableError,95 exceptions.GPTRouterTooManyRequestsError,96 ]97 return create_base_retry_decorator(98 error_types=errors, max_retries=llm.max_retries, run_manager=run_manager99 )100 101 102def completion_with_retry(103 llm: GPTRouter,104 models_priority_list: List[GPTRouterModel],105 run_manager: Optional[CallbackManagerForLLMRun] = None,106 **kwargs: Any,107) -> Union[GenerationResponse, Generator[ChunkedGenerationResponse, None, None]]:108 """Use tenacity to retry the completion call."""109 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)110 111 @retry_decorator112 def _completion_with_retry(**kwargs: Any) -> Any:113 ordered_generation_requests = get_ordered_generation_requests(114 models_priority_list, **kwargs115 )116 return llm.client.generate(117 ordered_generation_requests=ordered_generation_requests,118 is_stream=kwargs.get("stream", False),119 )120 121 return _completion_with_retry(**kwargs)122 123 124async def acompletion_with_retry(125 llm: GPTRouter,126 models_priority_list: List[GPTRouterModel],127 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,128 **kwargs: Any,129) -> Union[GenerationResponse, AsyncGenerator[ChunkedGenerationResponse, None]]:130 """Use tenacity to retry the async completion call."""131 132 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)133 134 @retry_decorator135 async def _completion_with_retry(**kwargs: Any) -> Any:136 ordered_generation_requests = get_ordered_generation_requests(137 models_priority_list, **kwargs138 )139 return await llm.client.agenerate(140 ordered_generation_requests=ordered_generation_requests,141 is_stream=kwargs.get("stream", False),142 )143 144 return await _completion_with_retry(**kwargs)145 146 147class GPTRouter(BaseChatModel):148 """GPTRouter by Writesonic Inc.149 150 For more information, see https://gpt-router.writesonic.com/docs151 """152 153 client: Any = Field(default=None, exclude=True) #: :meta private:154 models_priority_list: List[GPTRouterModel] = Field(min_length=1)155 gpt_router_api_base: str = Field(default="")156 """WriteSonic GPTRouter custom endpoint"""157 gpt_router_api_key: Optional[SecretStr] = None158 """WriteSonic GPTRouter API Key"""159 temperature: float = 0.7160 """What sampling temperature to use."""161 model_kwargs: Dict[str, Any] = Field(default_factory=dict)162 """Holds any model parameters valid for `create` call not explicitly specified."""163 max_retries: int = 4164 """Maximum number of retries to make when generating."""165 streaming: bool = False166 """Whether to stream the results or not."""167 n: int = 1168 """Number of chat completions to generate for each prompt."""169 max_tokens: int = 256170 171 @model_validator(mode="before")172 @classmethod173 def validate_environment(cls, values: Dict) -> Any:174 values["gpt_router_api_base"] = get_from_dict_or_env(175 values,176 "gpt_router_api_base",177 "GPT_ROUTER_API_BASE",178 DEFAULT_API_BASE_URL,179 )180 181 values["gpt_router_api_key"] = convert_to_secret_str(182 get_from_dict_or_env(183 values,184 "gpt_router_api_key",185 "GPT_ROUTER_API_KEY",186 )187 )188 return values189 190 @model_validator(mode="after")191 def post_init(self) -> Self:192 try:193 from gpt_router.client import GPTRouterClient194 195 except ImportError:196 raise GPTRouterException(197 "Could not import GPTRouter python package. "198 "Please install it with `pip install GPTRouter`."199 )200 201 gpt_router_client = GPTRouterClient(202 self.gpt_router_api_base,203 self.gpt_router_api_key.get_secret_value()204 if self.gpt_router_api_key205 else None,206 )207 self.client = gpt_router_client208 209 return self210 211 @property212 def lc_secrets(self) -> Dict[str, str]:213 return {"gpt_router_api_key": "GPT_ROUTER_API_KEY"}214 215 @property216 def lc_serializable(self) -> bool:217 return True218 219 @property220 def _llm_type(self) -> str:221 """Return type of chat model."""222 return "gpt-router-chat"223 224 @property225 def _identifying_params(self) -> Dict[str, Any]:226 """Get the identifying parameters."""227 return {228 **{"models_priority_list": self.models_priority_list},229 **self._default_params,230 }231 232 @property233 def _default_params(self) -> Dict[str, Any]:234 """Get the default parameters for calling GPTRouter API."""235 return {236 "max_tokens": self.max_tokens,237 "stream": self.streaming,238 "n": self.n,239 "temperature": self.temperature,240 **self.model_kwargs,241 }242 243 def _generate(244 self,245 messages: List[BaseMessage],246 stop: Optional[List[str]] = None,247 run_manager: Optional[CallbackManagerForLLMRun] = None,248 stream: Optional[bool] = None,249 **kwargs: Any,250 ) -> ChatResult:251 should_stream = stream if stream is not None else self.streaming252 if should_stream:253 stream_iter = self._stream(254 messages, stop=stop, run_manager=run_manager, **kwargs255 )256 return generate_from_stream(stream_iter)257 258 message_dicts, params = self._create_message_dicts(messages, stop)259 params = {**params, **kwargs, "stream": False}260 response = completion_with_retry(261 self,262 messages=message_dicts,263 models_priority_list=self.models_priority_list,264 run_manager=run_manager,265 **params,266 )267 return self._create_chat_result(response)268 269 async def _agenerate(270 self,271 messages: List[BaseMessage],272 stop: Optional[List[str]] = None,273 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,274 stream: Optional[bool] = None,275 **kwargs: Any,276 ) -> ChatResult:277 should_stream = stream if stream is not None else self.streaming278 if should_stream:279 stream_iter = self._astream(280 messages, stop=stop, run_manager=run_manager, **kwargs281 )282 return await agenerate_from_stream(stream_iter)283 284 message_dicts, params = self._create_message_dicts(messages, stop)285 params = {**params, **kwargs, "stream": False}286 response = await acompletion_with_retry(287 self,288 messages=message_dicts,289 models_priority_list=self.models_priority_list,290 run_manager=run_manager,291 **params,292 )293 return self._create_chat_result(response)294 295 def _create_chat_generation_chunk(296 self, data: Mapping[str, Any], default_chunk_class: Type[BaseMessageChunk]297 ) -> Tuple[ChatGenerationChunk, Type[BaseMessageChunk]]:298 chunk = _convert_delta_to_message_chunk(299 {"content": data.get("text", "")}, default_chunk_class300 )301 finish_reason = data.get("finish_reason")302 generation_info = (303 dict(finish_reason=finish_reason) if finish_reason is not None else None304 )305 default_chunk_class = chunk.__class__306 gen_chunk = ChatGenerationChunk(message=chunk, generation_info=generation_info)307 return gen_chunk, default_chunk_class308 309 def _stream(310 self,311 messages: List[BaseMessage],312 stop: Optional[List[str]] = None,313 run_manager: Optional[CallbackManagerForLLMRun] = None,314 **kwargs: Any,315 ) -> Iterator[ChatGenerationChunk]:316 message_dicts, params = self._create_message_dicts(messages, stop)317 params = {**params, **kwargs, "stream": True}318 319 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk320 generator_response = completion_with_retry(321 self,322 messages=message_dicts,323 models_priority_list=self.models_priority_list,324 run_manager=run_manager,325 **params,326 )327 for chunk in generator_response:328 if chunk.event != "update":329 continue330 331 chunk, default_chunk_class = self._create_chat_generation_chunk(332 chunk.data, default_chunk_class333 )334 335 if run_manager:336 run_manager.on_llm_new_token(337 token=str(chunk.message.content), chunk=chunk338 )339 340 yield chunk341 342 async def _astream(343 self,344 messages: List[BaseMessage],345 stop: Optional[List[str]] = None,346 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,347 **kwargs: Any,348 ) -> AsyncIterator[ChatGenerationChunk]:349 message_dicts, params = self._create_message_dicts(messages, stop)350 params = {**params, **kwargs, "stream": True}351 352 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk353 generator_response = acompletion_with_retry(354 self,355 messages=message_dicts,356 models_priority_list=self.models_priority_list,357 run_manager=run_manager,358 **params,359 )360 async for chunk in await generator_response:361 if chunk.event != "update":362 continue363 364 chunk, default_chunk_class = self._create_chat_generation_chunk(365 chunk.data, default_chunk_class366 )367 368 if run_manager:369 await run_manager.on_llm_new_token(370 token=str(chunk.message.content), chunk=chunk371 )372 373 yield chunk374 375 def _create_message_dicts(376 self, messages: List[BaseMessage], stop: Optional[List[str]]377 ) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]:378 params = self._default_params379 if stop is not None:380 if "stop" in params:381 raise ValueError("`stop` found in both the input and default params.")382 params["stop"] = stop383 message_dicts = [convert_message_to_dict(m) for m in messages]384 return message_dicts, params385 386 def _create_chat_result(self, response: GenerationResponse) -> ChatResult:387 generations = []388 for res in response.choices:389 message = convert_dict_to_message(390 {391 "role": "assistant",392 "content": res.text,393 }394 )395 gen = ChatGeneration(396 message=message,397 generation_info=dict(finish_reason=res.finish_reason),398 )399 generations.append(gen)400 llm_output = {"token_usage": response.meta, "model": response.model}401 return ChatResult(generations=generations, llm_output=llm_output)402 