Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
perplexity.py526 linesDownload Raw Back to chat_models
1"""Wrapper around Perplexity APIs."""2 3from __future__ import annotations4 5import logging6from operator import itemgetter7from typing import (8    Any,9    Dict,10    Iterator,11    List,12    Literal,13    Mapping,14    Optional,15    Tuple,16    Type,17    TypeVar,18    Union,19)20 21from langchain_core._api.deprecation import deprecated22from langchain_core.callbacks import CallbackManagerForLLMRun23from langchain_core.language_models import LanguageModelInput24from langchain_core.language_models.chat_models import (25    BaseChatModel,26    generate_from_stream,27)28from langchain_core.messages import (29    AIMessage,30    AIMessageChunk,31    BaseMessage,32    BaseMessageChunk,33    ChatMessage,34    ChatMessageChunk,35    FunctionMessageChunk,36    HumanMessage,37    HumanMessageChunk,38    SystemMessage,39    SystemMessageChunk,40    ToolMessageChunk,41)42from langchain_core.messages.ai import UsageMetadata43from langchain_core.output_parsers import JsonOutputParser, PydanticOutputParser44from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult45from langchain_core.runnables import Runnable, RunnableMap, RunnablePassthrough46from langchain_core.utils import from_env, get_pydantic_field_names47from langchain_core.utils.pydantic import (48    is_basemodel_subclass,49)50from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, model_validator51from typing_extensions import Self52 53_BM = TypeVar("_BM", bound=BaseModel)54_DictOrPydanticClass = Union[Dict[str, Any], Type[_BM], Type]55_DictOrPydantic = Union[Dict, _BM]56 57logger = logging.getLogger(__name__)58 59 60def _is_pydantic_class(obj: Any) -> bool:61    return isinstance(obj, type) and is_basemodel_subclass(obj)62 63 64def _create_usage_metadata(token_usage: dict) -> UsageMetadata:65    input_tokens = token_usage.get("prompt_tokens", 0)66    output_tokens = token_usage.get("completion_tokens", 0)67    total_tokens = token_usage.get("total_tokens", input_tokens + output_tokens)68    return UsageMetadata(69        input_tokens=input_tokens,70        output_tokens=output_tokens,71        total_tokens=total_tokens,72    )73 74 75@deprecated(76    since="0.3.21",77    removal="1.0",78    alternative_import="langchain_perplexity.ChatPerplexity",79)80class ChatPerplexity(BaseChatModel):81    """`Perplexity AI` Chat models API.82 83    Setup:84        To use, you should have the ``openai`` python package installed, and the85        environment variable ``PPLX_API_KEY`` set to your API key.86        Any parameters that are valid to be passed to the openai.create call87        can be passed in, even if not explicitly saved on this class.88 89        .. code-block:: bash90 91            pip install openai92            export PPLX_API_KEY=your_api_key93 94        Key init args - completion params:95            model: str96                Name of the model to use. e.g. "llama-3.1-sonar-small-128k-online"97            temperature: float98                Sampling temperature to use. Default is 0.799            max_tokens: Optional[int]100                Maximum number of tokens to generate.101            streaming: bool102                Whether to stream the results or not.103 104        Key init args - client params:105            pplx_api_key: Optional[str]106                API key for PerplexityChat API. Default is None.107            request_timeout: Optional[Union[float, Tuple[float, float]]]108                Timeout for requests to PerplexityChat completion API. Default is None.109            max_retries: int110                Maximum number of retries to make when generating.111 112        See full list of supported init args and their descriptions in the params section.113 114        Instantiate:115            .. code-block:: python116 117                from langchain_community.chat_models import ChatPerplexity118 119                llm = ChatPerplexity(120                    model="llama-3.1-sonar-small-128k-online",121                    temperature=0.7,122                )123 124        Invoke:125            .. code-block:: python126 127                messages = [128                    ("system", "You are a chatbot."),129                    ("user", "Hello!")130                ]131                llm.invoke(messages)132 133        Invoke with structured output:134            .. code-block:: python135 136                from pydantic import BaseModel137 138                class StructuredOutput(BaseModel):139                    role: str140                    content: str141 142                llm.with_structured_output(StructuredOutput)143                llm.invoke(messages)144 145        Invoke with perplexity-specific params:146            .. code-block:: python147 148                llm.invoke(messages, extra_body={"search_recency_filter": "week"})149 150        Stream:151            .. code-block:: python152 153                for chunk in llm.stream(messages):154                    print(chunk.content)155 156        Token usage:157            .. code-block:: python158 159                response = llm.invoke(messages)160                response.usage_metadata161 162        Response metadata:163            .. code-block:: python164 165                response = llm.invoke(messages)166                response.response_metadata167 168    """  # noqa: E501169 170    client: Any = None  #: :meta private:171    model: str = "llama-3.1-sonar-small-128k-online"172    """Model name."""173    temperature: float = 0.7174    """What sampling temperature to use."""175    model_kwargs: Dict[str, Any] = Field(default_factory=dict)176    """Holds any model parameters valid for `create` call not explicitly specified."""177    pplx_api_key: Optional[str] = Field(178        default_factory=from_env("PPLX_API_KEY", default=None), alias="api_key"179    )180    """Base URL path for API requests,181    leave blank if not using a proxy or service emulator."""182    request_timeout: Optional[Union[float, Tuple[float, float]]] = Field(183        None, alias="timeout"184    )185    """Timeout for requests to PerplexityChat completion API. Default is None."""186    max_retries: int = 6187    """Maximum number of retries to make when generating."""188    streaming: bool = False189    """Whether to stream the results or not."""190    max_tokens: Optional[int] = None191    """Maximum number of tokens to generate."""192 193    model_config = ConfigDict(194        populate_by_name=True,195    )196 197    @property198    def lc_secrets(self) -> Dict[str, str]:199        return {"pplx_api_key": "PPLX_API_KEY"}200 201    @model_validator(mode="before")202    @classmethod203    def build_extra(cls, values: Dict[str, Any]) -> Any:204        """Build extra kwargs from additional params that were passed in."""205        all_required_field_names = get_pydantic_field_names(cls)206        extra = values.get("model_kwargs", {})207        for field_name in list(values):208            if field_name in extra:209                raise ValueError(f"Found {field_name} supplied twice.")210            if field_name not in all_required_field_names:211                logger.warning(212                    f"""WARNING! {field_name} is not a default parameter.213                    {field_name} was transferred to model_kwargs.214                    Please confirm that {field_name} is what you intended."""215                )216                extra[field_name] = values.pop(field_name)217 218        invalid_model_kwargs = all_required_field_names.intersection(extra.keys())219        if invalid_model_kwargs:220            raise ValueError(221                f"Parameters {invalid_model_kwargs} should be specified explicitly. "222                f"Instead they were passed in as part of `model_kwargs` parameter."223            )224 225        values["model_kwargs"] = extra226        return values227 228    @model_validator(mode="after")229    def validate_environment(self) -> Self:230        """Validate that api key and python package exists in environment."""231        try:232            import openai233        except ImportError:234            raise ImportError(235                "Could not import openai python package. "236                "Please install it with `pip install openai`."237            )238        try:239            self.client = openai.OpenAI(240                api_key=self.pplx_api_key, base_url="https://api.perplexity.ai"241            )242        except AttributeError:243            raise ValueError(244                "`openai` has no `ChatCompletion` attribute, this is likely "245                "due to an old version of the openai package. Try upgrading it "246                "with `pip install --upgrade openai`."247            )248        return self249 250    @property251    def _default_params(self) -> Dict[str, Any]:252        """Get the default parameters for calling PerplexityChat API."""253        return {254            "max_tokens": self.max_tokens,255            "stream": self.streaming,256            "temperature": self.temperature,257            **self.model_kwargs,258        }259 260    def _convert_message_to_dict(self, message: BaseMessage) -> Dict[str, Any]:261        if isinstance(message, ChatMessage):262            message_dict = {"role": message.role, "content": message.content}263        elif isinstance(message, SystemMessage):264            message_dict = {"role": "system", "content": message.content}265        elif isinstance(message, HumanMessage):266            message_dict = {"role": "user", "content": message.content}267        elif isinstance(message, AIMessage):268            message_dict = {"role": "assistant", "content": message.content}269        else:270            raise TypeError(f"Got unknown type {message}")271        return message_dict272 273    def _create_message_dicts(274        self, messages: List[BaseMessage], stop: Optional[List[str]]275    ) -> Tuple[List[Dict[str, Any]], Dict[str, Any]]:276        params = dict(self._invocation_params)277        if stop is not None:278            if "stop" in params:279                raise ValueError("`stop` found in both the input and default params.")280            params["stop"] = stop281        message_dicts = [self._convert_message_to_dict(m) for m in messages]282        return message_dicts, params283 284    def _convert_delta_to_message_chunk(285        self, _dict: Mapping[str, Any], default_class: Type[BaseMessageChunk]286    ) -> BaseMessageChunk:287        role = _dict.get("role")288        content = _dict.get("content") or ""289        additional_kwargs: Dict = {}290        if _dict.get("function_call"):291            function_call = dict(_dict["function_call"])292            if "name" in function_call and function_call["name"] is None:293                function_call["name"] = ""294            additional_kwargs["function_call"] = function_call295        if _dict.get("tool_calls"):296            additional_kwargs["tool_calls"] = _dict["tool_calls"]297 298        if role == "user" or default_class == HumanMessageChunk:299            return HumanMessageChunk(content=content)300        elif role == "assistant" or default_class == AIMessageChunk:301            return AIMessageChunk(content=content, additional_kwargs=additional_kwargs)302        elif role == "system" or default_class == SystemMessageChunk:303            return SystemMessageChunk(content=content)304        elif role == "function" or default_class == FunctionMessageChunk:305            return FunctionMessageChunk(content=content, name=_dict["name"])306        elif role == "tool" or default_class == ToolMessageChunk:307            return ToolMessageChunk(content=content, tool_call_id=_dict["tool_call_id"])308        elif role or default_class == ChatMessageChunk:309            return ChatMessageChunk(content=content, role=role)  # type: ignore[arg-type]310        else:311            return default_class(content=content)  # type: ignore[call-arg]312 313    def _stream(314        self,315        messages: List[BaseMessage],316        stop: Optional[List[str]] = None,317        run_manager: Optional[CallbackManagerForLLMRun] = None,318        **kwargs: Any,319    ) -> Iterator[ChatGenerationChunk]:320        message_dicts, params = self._create_message_dicts(messages, stop)321        params = {**params, **kwargs}322        default_chunk_class = AIMessageChunk323        params.pop("stream", None)324        if stop:325            params["stop_sequences"] = stop326        stream_resp = self.client.chat.completions.create(327            messages=message_dicts, stream=True, **params328        )329        first_chunk = True330        prev_total_usage: Optional[UsageMetadata] = None331        for chunk in stream_resp:332            if not isinstance(chunk, dict):333                chunk = chunk.dict()334            # Collect standard usage metadata (transform from aggregate to delta)335            if total_usage := chunk.get("usage"):336                lc_total_usage = _create_usage_metadata(total_usage)337                if prev_total_usage:338                    usage_metadata: Optional[UsageMetadata] = {339                        "input_tokens": lc_total_usage["input_tokens"]340                        - prev_total_usage["input_tokens"],341                        "output_tokens": lc_total_usage["output_tokens"]342                        - prev_total_usage["output_tokens"],343                        "total_tokens": lc_total_usage["total_tokens"]344                        - prev_total_usage["total_tokens"],345                    }346                else:347                    usage_metadata = lc_total_usage348                prev_total_usage = lc_total_usage349            else:350                usage_metadata = None351            if len(chunk["choices"]) == 0:352                continue353            choice = chunk["choices"][0]354 355            additional_kwargs = {}356            if first_chunk:357                additional_kwargs["citations"] = chunk.get("citations", [])358                for attr in ["images", "related_questions"]:359                    if attr in chunk:360                        additional_kwargs[attr] = chunk[attr]361 362            chunk = self._convert_delta_to_message_chunk(363                choice["delta"], default_chunk_class364            )365 366            if isinstance(chunk, AIMessageChunk) and usage_metadata:367                chunk.usage_metadata = usage_metadata368 369            if first_chunk:370                chunk.additional_kwargs |= additional_kwargs371                first_chunk = False372 373            finish_reason = choice.get("finish_reason")374            generation_info = (375                dict(finish_reason=finish_reason) if finish_reason is not None else None376            )377            default_chunk_class = chunk.__class__378            chunk = ChatGenerationChunk(message=chunk, generation_info=generation_info)379            if run_manager:380                run_manager.on_llm_new_token(chunk.text, chunk=chunk)381            yield chunk382 383    def _generate(384        self,385        messages: List[BaseMessage],386        stop: Optional[List[str]] = None,387        run_manager: Optional[CallbackManagerForLLMRun] = None,388        **kwargs: Any,389    ) -> ChatResult:390        if self.streaming:391            stream_iter = self._stream(392                messages, stop=stop, run_manager=run_manager, **kwargs393            )394            if stream_iter:395                return generate_from_stream(stream_iter)396        message_dicts, params = self._create_message_dicts(messages, stop)397        params = {**params, **kwargs}398        response = self.client.chat.completions.create(messages=message_dicts, **params)399        if usage := getattr(response, "usage", None):400            usage_metadata = _create_usage_metadata(usage.model_dump())401        else:402            usage_metadata = None403 404        additional_kwargs = {"citations": response.citations}405        for attr in ["images", "related_questions"]:406            if hasattr(response, attr):407                additional_kwargs[attr] = getattr(response, attr)408 409        message = AIMessage(410            content=response.choices[0].message.content,411            additional_kwargs=additional_kwargs,412            usage_metadata=usage_metadata,413        )414        return ChatResult(generations=[ChatGeneration(message=message)])415 416    @property417    def _invocation_params(self) -> Mapping[str, Any]:418        """Get the parameters used to invoke the model."""419        pplx_creds: Dict[str, Any] = {420            "model": self.model,421        }422        return {**pplx_creds, **self._default_params}423 424    @property425    def _llm_type(self) -> str:426        """Return type of chat model."""427        return "perplexitychat"428 429    def with_structured_output(430        self,431        schema: Optional[_DictOrPydanticClass] = None,432        *,433        method: Literal["json_schema"] = "json_schema",434        include_raw: bool = False,435        strict: Optional[bool] = None,436        **kwargs: Any,437    ) -> Runnable[LanguageModelInput, _DictOrPydantic]:438        """Model wrapper that returns outputs formatted to match the given schema for Preplexity.439        Currently, Preplexity only supports "json_schema" method for structured output440        as per their official documentation: https://docs.perplexity.ai/guides/structured-outputs441 442        Args:443            schema:444                The output schema. Can be passed in as:445 446                - a JSON Schema,447                - a TypedDict class,448                - or a Pydantic class449 450            method: The method for steering model generation, currently only support:451 452                - "json_schema": Use the JSON Schema to parse the model output453 454 455            include_raw:456                If False then only the parsed structured output is returned. If457                an error occurs during model output parsing it will be raised. If True458                then both the raw model response (a BaseMessage) and the parsed model459                response will be returned. If an error occurs during output parsing it460                will be caught and returned as well. The final output is always a dict461                with keys "raw", "parsed", and "parsing_error".462 463            kwargs: Additional keyword args aren't supported.464 465        Returns:466            A Runnable that takes same inputs as a :class:`langchain_core.language_models.chat.BaseChatModel`.467 468            | If ``include_raw`` is False and ``schema`` is a Pydantic class, Runnable outputs an instance of ``schema`` (i.e., a Pydantic object). Otherwise, if ``include_raw`` is False then Runnable outputs a dict.469 470            | If ``include_raw`` is True, then Runnable outputs a dict with keys:471 472            - "raw": BaseMessage473            - "parsed": None if there was a parsing error, otherwise the type depends on the ``schema`` as described above.474            - "parsing_error": Optional[BaseException]475 476        """  # noqa: E501477        if method in ("function_calling", "json_mode"):478            method = "json_schema"479        if method == "json_schema":480            if schema is None:481                raise ValueError(482                    "schema must be specified when method is not 'json_schema'. "483                    "Received None."484                )485            is_pydantic_schema = _is_pydantic_class(schema)486            if is_pydantic_schema and hasattr(487                schema, "model_json_schema"488            ):  # accounting for pydantic v1 and v2489                response_format = schema.model_json_schema()490            elif is_pydantic_schema:491                response_format = schema.schema()  # type: ignore[union-attr]492            elif isinstance(schema, dict):493                response_format = schema494            elif type(schema).__name__ == "_TypedDictMeta":495                adapter = TypeAdapter(schema)  # if use passes typeddict496                response_format = adapter.json_schema()497 498            llm = self.bind(499                response_format={500                    "type": "json_schema",501                    "json_schema": {"schema": response_format},502                }503            )504            output_parser = (505                PydanticOutputParser(pydantic_object=schema)  # type: ignore[arg-type]506                if is_pydantic_schema507                else JsonOutputParser()508            )509        else:510            raise ValueError(511                f"Unrecognized method argument. Expected 'json_schema' Received:\512                    '{method}'"513            )514 515        if include_raw:516            parser_assign = RunnablePassthrough.assign(517                parsed=itemgetter("raw") | output_parser, parsing_error=lambda _: None518            )519            parser_none = RunnablePassthrough.assign(parsed=lambda _: None)520            parser_with_fallback = parser_assign.with_fallbacks(521                [parser_none], exception_key="parsing_error"522            )523            return RunnableMap(raw=llm) | parser_with_fallback524        else:525            return llm | output_parser526 
codekingpro/portable-devtools · Team Ai