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