codekingpro/portable-devtools
114k
1"""Wrapper around Prem's Chat API."""2 3from __future__ import annotations4 5import logging6import warnings7from typing import (8 TYPE_CHECKING,9 Any,10 Callable,11 Dict,12 Iterator,13 List,14 Optional,15 Sequence,16 Tuple,17 Type,18 Union,19)20 21from langchain_core.callbacks import (22 CallbackManagerForLLMRun,23)24from langchain_core.language_models import LanguageModelInput25from langchain_core.language_models.chat_models import BaseChatModel26from langchain_core.language_models.llms import create_base_retry_decorator27from langchain_core.messages import (28 AIMessage,29 AIMessageChunk,30 BaseMessage,31 BaseMessageChunk,32 ChatMessage,33 ChatMessageChunk,34 HumanMessage,35 HumanMessageChunk,36 SystemMessage,37 SystemMessageChunk,38 ToolMessage,39)40from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult41from langchain_core.runnables import Runnable42from langchain_core.tools import BaseTool43from langchain_core.utils import get_from_dict_or_env, pre_init44from langchain_core.utils.function_calling import convert_to_openai_tool45from pydantic import (46 BaseModel,47 ConfigDict,48 Field,49 SecretStr,50)51 52if TYPE_CHECKING:53 from premai.api.chat_completions.v1_chat_completions_create import (54 ChatCompletionResponseStream,55 )56 from premai.models.chat_completion_response import ChatCompletionResponse57 58logger = logging.getLogger(__name__)59 60TOOL_PROMPT_HEADER = """61Given the set of tools you used and the response, provide the final answer\n62"""63 64INTERMEDIATE_TOOL_RESULT_TEMPLATE = """65{json}66"""67 68SINGLE_TOOL_PROMPT_TEMPLATE = """69tool id: {tool_id}70tool_response: {tool_response}71"""72 73 74class ChatPremAPIError(Exception):75 """Error with the `PremAI` API."""76 77 78def _truncate_at_stop_tokens(79 text: str,80 stop: Optional[List[str]],81) -> str:82 """Truncates text at the earliest stop token found."""83 if stop is None:84 return text85 86 for stop_token in stop:87 stop_token_idx = text.find(stop_token)88 if stop_token_idx != -1:89 text = text[:stop_token_idx]90 return text91 92 93def _response_to_result(94 response: ChatCompletionResponse,95 stop: Optional[List[str]],96) -> ChatResult:97 """Converts a Prem API response into a LangChain result"""98 99 if not response.choices:100 raise ChatPremAPIError("ChatResponse must have at least one candidate")101 generations: List[ChatGeneration] = []102 for choice in response.choices:103 role = choice.message.role104 if role is None:105 raise ChatPremAPIError(f"ChatResponse {choice} must have a role.")106 107 # If content is None then it will be replaced by ""108 content = _truncate_at_stop_tokens(text=choice.message.content or "", stop=stop)109 if content is None:110 raise ChatPremAPIError(f"ChatResponse must have a content: {content}")111 112 if role == "assistant":113 tool_calls = choice.message["tool_calls"]114 if tool_calls is None:115 tools = []116 else:117 tools = [118 {119 "id": tool_call["id"],120 "name": tool_call["function"]["name"],121 "args": tool_call["function"]["arguments"],122 }123 for tool_call in tool_calls124 ]125 generations.append(126 ChatGeneration(127 text=content, message=AIMessage(content=content, tool_calls=tools)128 )129 )130 elif role == "user":131 generations.append(132 ChatGeneration(text=content, message=HumanMessage(content=content))133 )134 else:135 generations.append(136 ChatGeneration(137 text=content, message=ChatMessage(role=role, content=content)138 )139 )140 141 if response.document_chunks is not None:142 return ChatResult(143 generations=generations,144 llm_output={145 "document_chunks": [146 chunk.to_dict() for chunk in response.document_chunks147 ]148 },149 )150 else:151 return ChatResult(generations=generations, llm_output={"document_chunks": None})152 153 154def _convert_delta_response_to_message_chunk(155 response: ChatCompletionResponseStream, default_class: Type[BaseMessageChunk]156) -> Tuple[157 Union[BaseMessageChunk, HumanMessageChunk, AIMessageChunk, SystemMessageChunk],158 Optional[str],159]:160 """Converts delta response to message chunk"""161 _delta = response.choices[0].delta162 role = _delta.get("role", "")163 content = _delta.get("content", "")164 additional_kwargs: Dict = {}165 finish_reasons: Optional[str] = response.choices[0].finish_reason166 167 if role == "user" or default_class == HumanMessageChunk:168 return HumanMessageChunk(content=content), finish_reasons169 elif role == "assistant" or default_class == AIMessageChunk:170 return (171 AIMessageChunk(content=content, additional_kwargs=additional_kwargs),172 finish_reasons,173 )174 elif role == "system" or default_class == SystemMessageChunk:175 return SystemMessageChunk(content=content), finish_reasons176 elif role or default_class == ChatMessageChunk:177 return ChatMessageChunk(content=content, role=role), finish_reasons178 else:179 return default_class(content=content), finish_reasons # type: ignore[call-arg]180 181 182def _messages_to_prompt_dict(183 input_messages: List[BaseMessage],184 template_id: Optional[str] = None,185) -> Tuple[Optional[str], List[Dict[str, Any]]]:186 """Converts a list of LangChain Messages into a simple dict187 which is the message structure in Prem"""188 189 system_prompt: Optional[str] = None190 examples_and_messages: List[Dict[str, Any]] = []191 192 for input_msg in input_messages:193 if isinstance(input_msg, SystemMessage):194 system_prompt = str(input_msg.content)195 196 elif isinstance(input_msg, HumanMessage):197 if template_id is None:198 examples_and_messages.append(199 {200 "role": "user",201 "content": str(input_msg.content),202 }203 )204 else:205 params: Dict[str, str] = {}206 assert (input_msg.id is not None) and (input_msg.id != ""), ValueError(207 "When using prompt template there should be id associated ",208 "with each HumanMessage",209 )210 params[str(input_msg.id)] = str(input_msg.content)211 examples_and_messages.append(212 {213 "role": "user",214 "template_id": template_id,215 "params": params,216 }217 )218 elif isinstance(input_msg, AIMessage):219 if input_msg.tool_calls is None or len(input_msg.tool_calls) == 0:220 examples_and_messages.append(221 {222 "role": "assistant",223 "content": str(input_msg.content),224 }225 )226 else:227 ai_msg_to_json = {228 "id": input_msg.id,229 "content": input_msg.content,230 "response_metadata": input_msg.response_metadata,231 "tool_calls": input_msg.tool_calls,232 }233 examples_and_messages.append(234 {235 "role": "assistant",236 "content": INTERMEDIATE_TOOL_RESULT_TEMPLATE.format(237 json=ai_msg_to_json,238 ),239 }240 )241 elif isinstance(input_msg, ToolMessage):242 pass243 244 else:245 raise ChatPremAPIError("No such role explicitly exists")246 247 # do a separate search for tool calls248 tool_prompt = ""249 for input_msg in input_messages:250 if isinstance(input_msg, ToolMessage):251 tool_id = input_msg.tool_call_id252 tool_result = input_msg.content253 tool_prompt += SINGLE_TOOL_PROMPT_TEMPLATE.format(254 tool_id=tool_id, tool_response=tool_result255 )256 if tool_prompt != "":257 prompt = TOOL_PROMPT_HEADER258 prompt += tool_prompt259 examples_and_messages.append({"role": "user", "content": prompt})260 261 return system_prompt, examples_and_messages262 263 264class ChatPremAI(BaseChatModel, BaseModel):265 """PremAI Chat models.266 267 To use, you will need to have an API key. You can find your existing API Key268 or generate a new one here: https://app.premai.io/api_keys/269 """270 271 # TODO: Need to add the default parameters through prem-sdk here272 273 project_id: int274 """The project ID in which the experiments or deployments are carried out. 275 You can find all your projects here: https://app.premai.io/projects/"""276 premai_api_key: Optional[SecretStr] = Field(default=None, alias="api_key")277 """Prem AI API Key. Get it here: https://app.premai.io/api_keys/"""278 279 model: Optional[str] = Field(default=None, alias="model_name")280 """Name of the model. This is an optional parameter. 281 The default model is the one deployed from Prem's LaunchPad: https://app.premai.io/projects/8/launchpad282 If model name is other than default model then it will override the calls 283 from the model deployed from launchpad."""284 285 session_id: Optional[str] = None286 """The ID of the session to use. It helps to track the chat history."""287 288 temperature: Optional[float] = Field(default=None)289 """Model temperature. Value should be >= 0 and <= 1.0"""290 291 top_p: Optional[float] = None292 """top_p adjusts the number of choices for each predicted tokens based on293 cumulative probabilities. Value should be ranging between 0.0 and 1.0. 294 """295 296 max_tokens: Optional[int] = Field(default=None)297 298 """The maximum number of tokens to generate"""299 300 max_retries: int = Field(default=1)301 """Max number of retries to call the API"""302 303 system_prompt: Optional[str] = ""304 """Acts like a default instruction that helps the LLM act or generate 305 in a specific way.This is an Optional Parameter. By default the 306 system prompt would be using Prem's Launchpad models system prompt. 307 Changing the system prompt would override the default system prompt.308 """309 310 repositories: Optional[dict] = None311 """Add valid repository ids. This will be overriding existing connected 312 repositories (if any) and will use RAG with the connected repos. 313 """314 315 streaming: Optional[bool] = False316 """Whether to stream the responses or not."""317 318 client: Any = None319 320 model_config = ConfigDict(321 populate_by_name=True,322 arbitrary_types_allowed=True,323 extra="forbid",324 )325 326 @pre_init327 def validate_environments(cls, values: Dict) -> Dict:328 """Validate that the package is installed and that the API token is valid"""329 try:330 from premai import Prem331 except ImportError as error:332 raise ImportError(333 "Could not import Prem Python package."334 "Please install it with: `pip install premai`"335 ) from error336 337 try:338 premai_api_key: Union[str, SecretStr] = get_from_dict_or_env(339 values, "premai_api_key", "PREMAI_API_KEY"340 )341 values["client"] = Prem(342 api_key=premai_api_key343 if isinstance(premai_api_key, str)344 else premai_api_key._secret_value345 )346 except Exception as error:347 raise ValueError("Your API Key is incorrect. Please try again.") from error348 return values349 350 @property351 def _llm_type(self) -> str:352 return "premai"353 354 @property355 def _default_params(self) -> Dict[str, Any]:356 return {357 "model": self.model,358 "system_prompt": self.system_prompt,359 "temperature": self.temperature,360 "max_tokens": self.max_tokens,361 "repositories": self.repositories,362 }363 364 def _get_all_kwargs(self, **kwargs: Any) -> Dict[str, Any]:365 kwargs_to_ignore = [366 "top_p",367 "frequency_penalty",368 "presence_penalty",369 "logit_bias",370 "stop",371 "seed",372 ]373 keys_to_remove = []374 375 for key in kwargs:376 if key in kwargs_to_ignore:377 warnings.warn(f"WARNING: Parameter {key} is not supported in kwargs.")378 keys_to_remove.append(key)379 380 for key in keys_to_remove:381 kwargs.pop(key)382 383 all_kwargs = {**self._default_params, **kwargs}384 for key in list(self._default_params.keys()):385 if all_kwargs.get(key) is None or all_kwargs.get(key) == "":386 all_kwargs.pop(key, None)387 return all_kwargs388 389 def _generate(390 self,391 messages: List[BaseMessage],392 stop: Optional[List[str]] = None,393 run_manager: Optional[CallbackManagerForLLMRun] = None,394 **kwargs: Any,395 ) -> ChatResult:396 if "template_id" in kwargs:397 system_prompt, messages_to_pass = _messages_to_prompt_dict(398 messages, template_id=kwargs["template_id"]399 )400 else:401 system_prompt, messages_to_pass = _messages_to_prompt_dict(messages)402 403 if system_prompt is not None and system_prompt != "":404 kwargs["system_prompt"] = system_prompt405 406 all_kwargs = self._get_all_kwargs(**kwargs)407 response = chat_with_retry(408 self,409 project_id=self.project_id,410 messages=messages_to_pass,411 stream=False,412 run_manager=run_manager,413 **all_kwargs,414 )415 416 return _response_to_result(response=response, stop=stop)417 418 def _stream(419 self,420 messages: List[BaseMessage],421 stop: Optional[List[str]] = None,422 run_manager: Optional[CallbackManagerForLLMRun] = None,423 **kwargs: Any,424 ) -> Iterator[ChatGenerationChunk]:425 if "template_id" in kwargs:426 system_prompt, messages_to_pass = _messages_to_prompt_dict(427 messages, template_id=kwargs["template_id"]428 )429 else:430 system_prompt, messages_to_pass = _messages_to_prompt_dict(messages)431 432 if stop is not None:433 logger.warning("stop is not supported in langchain streaming")434 435 if "system_prompt" not in kwargs:436 if system_prompt is not None and system_prompt != "":437 kwargs["system_prompt"] = system_prompt438 439 all_kwargs = self._get_all_kwargs(**kwargs)440 441 default_chunk_class = AIMessageChunk442 443 for streamed_response in chat_with_retry(444 self,445 project_id=self.project_id,446 messages=messages_to_pass,447 stream=True,448 run_manager=run_manager,449 **all_kwargs,450 ):451 try:452 chunk, finish_reason = _convert_delta_response_to_message_chunk(453 response=streamed_response, default_class=default_chunk_class454 )455 generation_info = (456 dict(finish_reason=finish_reason)457 if finish_reason is not None458 else None459 )460 cg_chunk = ChatGenerationChunk(461 message=chunk, generation_info=generation_info462 )463 if run_manager:464 run_manager.on_llm_new_token(cg_chunk.text, chunk=cg_chunk)465 yield cg_chunk466 except Exception as _:467 continue468 469 def bind_tools(470 self,471 tools: Sequence[Union[Dict[str, Any], Type[BaseModel], Callable, BaseTool]],472 **kwargs: Any,473 ) -> Runnable[LanguageModelInput, AIMessage]:474 formatted_tools = [convert_to_openai_tool(tool) for tool in tools]475 return super().bind(tools=formatted_tools, **kwargs)476 477 478def create_prem_retry_decorator(479 llm: ChatPremAI,480 *,481 max_retries: int = 1,482 run_manager: Optional[Union[CallbackManagerForLLMRun]] = None,483) -> Callable[[Any], Any]:484 """Create a retry decorator for PremAI API errors."""485 import premai.models486 487 errors = [488 premai.models.api_response_validation_error.APIResponseValidationError,489 premai.models.conflict_error.ConflictError,490 premai.models.model_not_found_error.ModelNotFoundError,491 premai.models.permission_denied_error.PermissionDeniedError,492 premai.models.provider_api_connection_error.ProviderAPIConnectionError,493 premai.models.provider_api_status_error.ProviderAPIStatusError,494 premai.models.provider_api_timeout_error.ProviderAPITimeoutError,495 premai.models.provider_internal_server_error.ProviderInternalServerError,496 premai.models.provider_not_found_error.ProviderNotFoundError,497 premai.models.rate_limit_error.RateLimitError,498 premai.models.unprocessable_entity_error.UnprocessableEntityError,499 premai.models.validation_error.ValidationError,500 ]501 502 decorator = create_base_retry_decorator(503 error_types=errors, max_retries=max_retries, run_manager=run_manager504 )505 return decorator506 507 508def chat_with_retry(509 llm: ChatPremAI,510 project_id: int,511 messages: List[dict],512 stream: bool = False,513 run_manager: Optional[CallbackManagerForLLMRun] = None,514 **kwargs: Any,515) -> Any:516 """Using tenacity for retry in completion call"""517 retry_decorator = create_prem_retry_decorator(518 llm, max_retries=llm.max_retries, run_manager=run_manager519 )520 521 @retry_decorator522 def _completion_with_retry(523 project_id: int,524 messages: List[dict],525 stream: Optional[bool] = False,526 **kwargs: Any,527 ) -> Any:528 response = llm.client.chat.completions.create(529 project_id=project_id,530 messages=messages,531 stream=stream,532 **kwargs,533 )534 return response535 536 return _completion_with_retry(537 project_id=project_id,538 messages=messages,539 stream=stream,540 **kwargs,541 )542 