codekingpro/portable-devtools
114k
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 