Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
maritalk.py375 linesDownload Raw Back to chat_models
1import json2from http import HTTPStatus3from typing import Any, AsyncIterator, Dict, Iterator, List, Optional, Union4 5import requests6from langchain_core.callbacks import (7    AsyncCallbackManagerForLLMRun,8    CallbackManagerForLLMRun,9)10from langchain_core.language_models.chat_models import BaseChatModel11from langchain_core.messages import (12    AIMessage,13    AIMessageChunk,14    BaseMessage,15    HumanMessage,16    SystemMessage,17)18from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult19from pydantic import Field20from requests import Response21from requests.exceptions import HTTPError22 23 24class MaritalkHTTPError(HTTPError):25    def __init__(self, request_obj: Response) -> None:26        self.request_obj = request_obj27        try:28            response_json = request_obj.json()29            if "detail" in response_json:30                api_message = response_json["detail"]31            elif "message" in response_json:32                api_message = response_json["message"]33            else:34                api_message = response_json35        except Exception:36            api_message = request_obj.text37 38        self.message = api_message39        self.status_code = request_obj.status_code40 41    def __str__(self) -> str:42        status_code_meaning = HTTPStatus(self.status_code).phrase43        formatted_message = f"HTTP Error: {self.status_code} - {status_code_meaning}"44        formatted_message += f"\nDetail: {self.message}"45        return formatted_message46 47 48class ChatMaritalk(BaseChatModel):49    """`MariTalk` Chat models API.50 51    This class allows interacting with the MariTalk chatbot API.52    To use it, you must provide an API key either through the constructor.53 54    Example:55        .. code-block:: python56 57            from langchain_community.chat_models import ChatMaritalk58            chat = ChatMaritalk(api_key="your_api_key_here")59    """60 61    api_key: str62    """Your MariTalk API key."""63 64    model: str65    """Chose one of the available models: 66    - `sabia-2-medium`67    - `sabia-2-small`68    - `sabia-2-medium-2024-03-13`69    - `sabia-2-small-2024-03-13`70    - `maritalk-2024-01-08` (deprecated)"""71 72    temperature: float = Field(default=0.7, gt=0.0, lt=1.0)73    """Run inference with this temperature. 74    Must be in the closed interval [0.0, 1.0]."""75 76    max_tokens: int = Field(default=512, gt=0)77    """The maximum number of tokens to generate in the reply."""78 79    do_sample: bool = Field(default=True)80    """Whether or not to use sampling; use `True` to enable."""81 82    top_p: float = Field(default=0.95, gt=0.0, lt=1.0)83    """Nucleus sampling parameter controlling the size of 84    the probability mass considered for sampling."""85 86    @property87    def _llm_type(self) -> str:88        """Identifies the LLM type as 'maritalk'."""89        return "maritalk"90 91    def parse_messages_for_model(92        self, messages: List[BaseMessage]93    ) -> List[Dict[str, Union[str, List[Union[str, Dict[Any, Any]]]]]]:94        """95        Parses messages from LangChain's format to the format expected by96        the MariTalk API.97 98        Parameters:99            messages (List[BaseMessage]): A list of messages in LangChain100            format to be parsed.101 102        Returns:103            A list of messages formatted for the MariTalk API.104        """105        parsed_messages = []106 107        for message in messages:108            if isinstance(message, HumanMessage):109                role = "user"110            elif isinstance(message, AIMessage):111                role = "assistant"112            elif isinstance(message, SystemMessage):113                role = "system"114 115            parsed_messages.append({"role": role, "content": message.content})116        return parsed_messages117 118    def _call(119        self,120        messages: List[BaseMessage],121        stop: Optional[List[str]] = None,122        run_manager: Optional[CallbackManagerForLLMRun] = None,123        **kwargs: Any,124    ) -> str:125        """126        Sends the parsed messages to the MariTalk API and returns the generated127        response or an error message.128 129        This method makes an HTTP POST request to the MariTalk API with the130        provided messages and other parameters.131        If the request is successful and the API returns a response,132        this method returns a string containing the answer.133        If the request is rate-limited or encounters another error,134        it returns a string with the error message.135 136        Parameters:137            messages (List[BaseMessage]): Messages to send to the model.138            stop (Optional[List[str]]): Tokens that will signal the model139                to stop generating further tokens.140 141        Returns:142            str: If the API call is successful, returns the answer.143                 If an error occurs (e.g., rate limiting), returns a string144                 describing the error.145        """146        url = "https://chat.maritaca.ai/api/chat/inference"147        headers = {"authorization": f"Key {self.api_key}"}148        stopping_tokens = stop if stop is not None else []149 150        parsed_messages = self.parse_messages_for_model(messages)151 152        data = {153            "messages": parsed_messages,154            "model": self.model,155            "do_sample": self.do_sample,156            "max_tokens": self.max_tokens,157            "temperature": self.temperature,158            "top_p": self.top_p,159            "stopping_tokens": stopping_tokens,160            **kwargs,161        }162 163        response = requests.post(url, json=data, headers=headers)164 165        if response.ok:166            return response.json().get("answer", "No answer found")167        else:168            raise MaritalkHTTPError(response)169 170    async def _acall(171        self,172        messages: List[BaseMessage],173        stop: Optional[List[str]] = None,174        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,175        **kwargs: Any,176    ) -> str:177        """178        Asynchronously sends the parsed messages to the MariTalk API and returns179        the generated response or an error message.180 181        This method makes an HTTP POST request to the MariTalk API with the182        provided messages and other parameters using async I/O.183        If the request is successful and the API returns a response,184        this method returns a string containing the answer.185        If the request is rate-limited or encounters another error,186        it returns a string with the error message.187        """188        try:189            import httpx190 191            url = "https://chat.maritaca.ai/api/chat/inference"192            headers = {"authorization": f"Key {self.api_key}"}193            stopping_tokens = stop if stop is not None else []194 195            parsed_messages = self.parse_messages_for_model(messages)196 197            data = {198                "messages": parsed_messages,199                "model": self.model,200                "do_sample": self.do_sample,201                "max_tokens": self.max_tokens,202                "temperature": self.temperature,203                "top_p": self.top_p,204                "stopping_tokens": stopping_tokens,205                **kwargs,206            }207 208            async with httpx.AsyncClient() as client:209                response = await client.post(210                    url, json=data, headers=headers, timeout=None211                )212 213            if response.status_code == 200:214                return response.json().get("answer", "No answer found")215            else:216                raise MaritalkHTTPError(response)  # type: ignore[arg-type]217 218        except ImportError:219            raise ImportError(220                "Could not import httpx python package. "221                "Please install it with `pip install httpx`."222            )223 224    def _stream(225        self,226        messages: List[BaseMessage],227        stop: Optional[List[str]] = None,228        run_manager: Optional[CallbackManagerForLLMRun] = None,229        **kwargs: Any,230    ) -> Iterator[ChatGenerationChunk]:231        headers = {"Authorization": f"Key {self.api_key}"}232        stopping_tokens = stop if stop is not None else []233 234        parsed_messages = self.parse_messages_for_model(messages)235 236        data = {237            "messages": parsed_messages,238            "model": self.model,239            "do_sample": self.do_sample,240            "max_tokens": self.max_tokens,241            "temperature": self.temperature,242            "top_p": self.top_p,243            "stopping_tokens": stopping_tokens,244            "stream": True,245            **kwargs,246        }247 248        response = requests.post(249            "https://chat.maritaca.ai/api/chat/inference",250            data=json.dumps(data),251            headers=headers,252            stream=True,253        )254 255        if response.ok:256            for line in response.iter_lines():257                if line.startswith(b"data: "):258                    response_data = line.replace(b"data: ", b"").decode("utf-8")259                    if response_data:260                        parsed_data = json.loads(response_data)261                        if "text" in parsed_data:262                            delta = parsed_data["text"]263                            chunk = ChatGenerationChunk(264                                message=AIMessageChunk(content=delta)265                            )266                            if run_manager:267                                run_manager.on_llm_new_token(delta, chunk=chunk)268                            yield chunk269 270        else:271            raise MaritalkHTTPError(response)272 273    async def _astream(274        self,275        messages: List[BaseMessage],276        stop: Optional[List[str]] = None,277        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,278        **kwargs: Any,279    ) -> AsyncIterator[ChatGenerationChunk]:280        try:281            import httpx282 283            headers = {"Authorization": f"Key {self.api_key}"}284            stopping_tokens = stop if stop is not None else []285 286            parsed_messages = self.parse_messages_for_model(messages)287 288            data = {289                "messages": parsed_messages,290                "model": self.model,291                "do_sample": self.do_sample,292                "max_tokens": self.max_tokens,293                "temperature": self.temperature,294                "top_p": self.top_p,295                "stopping_tokens": stopping_tokens,296                "stream": True,297                **kwargs,298            }299 300            async with httpx.AsyncClient() as client:301                async with client.stream(302                    "POST",303                    "https://chat.maritaca.ai/api/chat/inference",304                    data=json.dumps(data),  # type: ignore[arg-type]305                    headers=headers,306                    timeout=None,307                ) as response:308                    if response.status_code == 200:309                        async for line in response.aiter_lines():310                            if line.startswith("data: "):311                                response_data = line.replace("data: ", "")312                                if response_data:313                                    parsed_data = json.loads(response_data)314                                    if "text" in parsed_data:315                                        delta = parsed_data["text"]316                                        chunk = ChatGenerationChunk(317                                            message=AIMessageChunk(content=delta)318                                        )319                                        if run_manager:320                                            await run_manager.on_llm_new_token(321                                                delta, chunk=chunk322                                            )323                                        yield chunk324 325                    else:326                        raise MaritalkHTTPError(response)  # type: ignore[arg-type]327 328        except ImportError:329            raise ImportError(330                "Could not import httpx python package. "331                "Please install it with `pip install httpx`."332            )333 334    def _generate(335        self,336        messages: List[BaseMessage],337        stop: Optional[List[str]] = None,338        run_manager: Optional[CallbackManagerForLLMRun] = None,339        **kwargs: Any,340    ) -> ChatResult:341        output_str = self._call(messages, stop=stop, run_manager=run_manager, **kwargs)342        message = AIMessage(content=output_str)343        generation = ChatGeneration(message=message)344        return ChatResult(generations=[generation])345 346    async def _agenerate(347        self,348        messages: List[BaseMessage],349        stop: Optional[List[str]] = None,350        run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,351        **kwargs: Any,352    ) -> ChatResult:353        output_str = await self._acall(354            messages, stop=stop, run_manager=run_manager, **kwargs355        )356        message = AIMessage(content=output_str)357        generation = ChatGeneration(message=message)358        return ChatResult(generations=[generation])359 360    @property361    def _identifying_params(self) -> Dict[str, Any]:362        """363        Identifies the key parameters of the chat model for logging364        or tracking purposes.365 366        Returns:367            A dictionary of the key configuration parameters.368        """369        return {370            "model": self.model,371            "temperature": self.temperature,372            "top_p": self.top_p,373            "max_tokens": self.max_tokens,374        }375 
codekingpro/portable-devtools · Team Ai