Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
gpt_router.py402 linesDownload Raw Back to chat_models
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 
codekingpro/portable-devtools · Team Ai