Team Ai
Datasetpublic

codekingpro/portable-devtools

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