Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
base.py479 linesDownload Raw Back to prompts
1"""Base class for prompt templates."""2 3from __future__ import annotations4 5import builtins  # noqa: TC0036import contextlib7import json8from abc import ABC, abstractmethod9from collections.abc import Mapping  # noqa: TC00310from functools import cached_property11from pathlib import Path12from typing import TYPE_CHECKING, Any, Generic, TypeVar, cast13 14import yaml15from pydantic import BaseModel, ConfigDict, Field, model_validator16from typing_extensions import Self, override17 18from langchain_core._api import deprecated19from langchain_core.exceptions import ErrorCode, create_message20from langchain_core.load import dumpd21from langchain_core.output_parsers.base import BaseOutputParser  # noqa: TC00122from langchain_core.prompt_values import (23    ChatPromptValueConcrete,24    PromptValue,25    StringPromptValue,26)27from langchain_core.runnables import RunnableConfig, RunnableSerializable28from langchain_core.runnables.config import ensure_config29from langchain_core.utils.pydantic import create_model_v230 31if TYPE_CHECKING:32    from collections.abc import Callable33 34    from langchain_core.documents import Document35 36 37FormatOutputType = TypeVar("FormatOutputType")38 39 40class BasePromptTemplate(41    RunnableSerializable[dict, PromptValue], ABC, Generic[FormatOutputType]42):43    """Base class for all prompt templates, returning a prompt."""44 45    input_variables: list[str]46    """A list of the names of the variables whose values are required as inputs to the47    prompt.48    """49 50    optional_variables: list[str] = Field(default=[])51    """A list of the names of the variables for placeholder or `MessagePlaceholder` that52    are optional.53 54    These variables are auto inferred from the prompt and user need not provide them.55    """56 57    input_types: builtins.dict[str, Any] = Field(default_factory=dict, exclude=True)58    """A dictionary of the types of the variables the prompt template expects.59 60    If not provided, all variables are assumed to be strings.61    """62 63    output_parser: BaseOutputParser | None = None64    """How to parse the output of calling an LLM on this formatted prompt."""65 66    partial_variables: Mapping[str, Any] = Field(default_factory=dict)67    """A dictionary of the partial variables the prompt template carries.68 69    Partial variables populate the template so that you don't need to pass them in every70    time you call the prompt.71    """72 73    metadata: builtins.dict[str, Any] | None = None74    """Metadata to be used for tracing."""75 76    tags: list[str] | None = None77    """Tags to be used for tracing."""78 79    @model_validator(mode="after")80    def validate_variable_names(self) -> Self:81        """Validate variable names do not include restricted names."""82        if "stop" in self.input_variables:83            msg = (84                "Cannot have an input variable named 'stop', as it is used internally,"85                " please rename."86            )87            raise ValueError(88                create_message(message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT)89            )90        if "stop" in self.partial_variables:91            msg = (92                "Cannot have an partial variable named 'stop', as it is used "93                "internally, please rename."94            )95            raise ValueError(96                create_message(message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT)97            )98 99        overall = set(self.input_variables).intersection(self.partial_variables)100        if overall:101            msg = f"Found overlapping input and partial variables: {overall}"102            raise ValueError(103                create_message(message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT)104            )105        return self106 107    @classmethod108    def get_lc_namespace(cls) -> list[str]:109        """Get the namespace of the LangChain object.110 111        Returns:112            `["langchain", "schema", "prompt_template"]`113        """114        return ["langchain", "schema", "prompt_template"]115 116    @classmethod117    def is_lc_serializable(cls) -> bool:118        """Return `True` as this class is serializable."""119        return True120 121    model_config = ConfigDict(122        arbitrary_types_allowed=True,123    )124 125    @cached_property126    def _serialized(self) -> dict[str, Any]:127        # self is always a Serializable object in this case, thus the result is128        # guaranteed to be a dict since dumpd uses the default callback, which uses129        # obj.to_json which always returns TypedDict subclasses130        return cast("dict[str, Any]", dumpd(self))131 132    @property133    @override134    def OutputType(self) -> Any:135        """Return the output type of the prompt."""136        return StringPromptValue | ChatPromptValueConcrete137 138    @override139    def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:140        """Get the input schema for the prompt.141 142        Args:143            config: Configuration for the prompt.144 145        Returns:146            The input schema for the prompt.147        """148        # This is correct, but pydantic typings/mypy don't think so.149        required_input_variables = {150            k: (self.input_types.get(k, str), ...) for k in self.input_variables151        }152        optional_input_variables = {153            k: (self.input_types.get(k, str), None) for k in self.optional_variables154        }155        return create_model_v2(156            "PromptInput",157            field_definitions={**required_input_variables, **optional_input_variables},158        )159 160    def _validate_input(self, inner_input: Any) -> dict:161        if not isinstance(inner_input, dict):162            if len(self.input_variables) == 1:163                var_name = self.input_variables[0]164                inner_input_ = {var_name: inner_input}165 166            else:167                msg = (168                    f"Expected mapping type as input to {self.__class__.__name__}. "169                    f"Received {type(inner_input)}."170                )171                raise TypeError(172                    create_message(173                        message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT174                    )175                )176        else:177            inner_input_ = inner_input178        missing = set(self.input_variables).difference(inner_input_)179        if missing:180            msg = (181                f"Input to {self.__class__.__name__} is missing variables {missing}. "182                f" Expected: {self.input_variables}"183                f" Received: {list(inner_input_.keys())}"184            )185            example_key = missing.pop()186            msg += (187                f"\nNote: if you intended {{{example_key}}} to be part of the string"188                " and not a variable, please escape it with double curly braces like: "189                f"'{{{{{example_key}}}}}'."190            )191            raise KeyError(192                create_message(message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT)193            )194        return inner_input_195 196    def _format_prompt_with_error_handling(self, inner_input: dict) -> PromptValue:197        inner_input_ = self._validate_input(inner_input)198        return self.format_prompt(**inner_input_)199 200    async def _aformat_prompt_with_error_handling(201        self, inner_input: dict202    ) -> PromptValue:203        inner_input_ = self._validate_input(inner_input)204        return await self.aformat_prompt(**inner_input_)205 206    @override207    def invoke(208        self, input: dict, config: RunnableConfig | None = None, **kwargs: Any209    ) -> PromptValue:210        """Invoke the prompt.211 212        Args:213            input: Input to the prompt.214            config: Configuration for the prompt.215 216        Returns:217            The output of the prompt.218        """219        config = ensure_config(config)220        if self.metadata:221            config["metadata"] = {**config["metadata"], **self.metadata}222        if self.tags:223            config["tags"] += self.tags224        return self._call_with_config(225            self._format_prompt_with_error_handling,226            input,227            config,228            run_type="prompt",229            serialized=self._serialized,230        )231 232    @override233    async def ainvoke(234        self, input: dict, config: RunnableConfig | None = None, **kwargs: Any235    ) -> PromptValue:236        """Async invoke the prompt.237 238        Args:239            input: Input to the prompt.240            config: Configuration for the prompt.241 242        Returns:243            The output of the prompt.244        """245        config = ensure_config(config)246        if self.metadata:247            config["metadata"].update(self.metadata)248        if self.tags:249            config["tags"].extend(self.tags)250        return await self._acall_with_config(251            self._aformat_prompt_with_error_handling,252            input,253            config,254            run_type="prompt",255            serialized=self._serialized,256        )257 258    @abstractmethod259    def format_prompt(self, **kwargs: Any) -> PromptValue:260        """Create `PromptValue`.261 262        Args:263            **kwargs: Any arguments to be passed to the prompt template.264 265        Returns:266            The output of the prompt.267        """268 269    async def aformat_prompt(self, **kwargs: Any) -> PromptValue:270        """Async create `PromptValue`.271 272        Args:273            **kwargs: Any arguments to be passed to the prompt template.274 275        Returns:276            The output of the prompt.277        """278        return self.format_prompt(**kwargs)279 280    def partial(self, **kwargs: str | Callable[[], str]) -> BasePromptTemplate:281        """Return a partial of the prompt template.282 283        Args:284            **kwargs: Partial variables to set.285 286        Returns:287            A partial of the prompt template.288        """289        prompt_dict = self.__dict__.copy()290        prompt_dict["input_variables"] = list(291            set(self.input_variables).difference(kwargs)292        )293        prompt_dict["partial_variables"] = {**self.partial_variables, **kwargs}294        return type(self)(**prompt_dict)295 296    def _merge_partial_and_user_variables(self, **kwargs: Any) -> dict[str, Any]:297        # Get partial params:298        partial_kwargs = {299            k: v if not callable(v) else v() for k, v in self.partial_variables.items()300        }301        return {**partial_kwargs, **kwargs}302 303    @abstractmethod304    def format(self, **kwargs: Any) -> FormatOutputType:305        """Format the prompt with the inputs.306 307        Args:308            **kwargs: Any arguments to be passed to the prompt template.309 310        Returns:311            A formatted string.312 313        Example:314            ```python315            prompt.format(variable1="foo")316            ```317        """318 319    async def aformat(self, **kwargs: Any) -> FormatOutputType:320        """Async format the prompt with the inputs.321 322        Args:323            **kwargs: Any arguments to be passed to the prompt template.324 325        Returns:326            A formatted string.327 328        Example:329            ```python330            await prompt.aformat(variable1="foo")331            ```332        """333        return self.format(**kwargs)334 335    @property336    def _prompt_type(self) -> str:337        """Return the prompt type key."""338        raise NotImplementedError339 340    def dict(self, **kwargs: Any) -> dict:341        """Return dictionary representation of prompt.342 343        Args:344            **kwargs: Any additional arguments to pass to the dictionary.345 346        Returns:347            Dictionary representation of the prompt.348        """349        prompt_dict = super().model_dump(**kwargs)350        with contextlib.suppress(NotImplementedError):351            prompt_dict["_type"] = self._prompt_type352        return prompt_dict353 354    @deprecated(355        since="1.2.21",356        removal="2.0.0",357        alternative="Use `dumpd`/`dumps` from `langchain_core.load` to serialize "358        "prompts and `load`/`loads` to deserialize them.",359    )360    def save(self, file_path: Path | str) -> None:361        """Save the prompt.362 363        Args:364            file_path: Path to directory to save prompt to.365 366        Raises:367            ValueError: If the prompt has partial variables.368            ValueError: If the file path is not json or yaml.369            NotImplementedError: If the prompt type is not implemented.370 371        Example:372            ```python373            prompt.save(file_path="path/prompt.yaml")374            ```375        """376        if self.partial_variables:377            msg = "Cannot save prompt with partial variables."378            raise ValueError(msg)379 380        # Fetch dictionary to save381        prompt_dict = self.dict()382        if "_type" not in prompt_dict:383            msg = f"Prompt {self} does not support saving."384            raise NotImplementedError(msg)385 386        # Convert file to Path object.387        save_path = Path(file_path)388 389        directory_path = save_path.parent390        directory_path.mkdir(parents=True, exist_ok=True)391 392        resolved_path = save_path.resolve()393        if resolved_path.suffix == ".json":394            with resolved_path.open("w", encoding="utf-8") as f:395                json.dump(prompt_dict, f, indent=4)396        elif resolved_path.suffix.endswith((".yaml", ".yml")):397            with resolved_path.open("w", encoding="utf-8") as f:398                yaml.dump(prompt_dict, f, default_flow_style=False)399        else:400            msg = f"{save_path} must be json or yaml"401            raise ValueError(msg)402 403 404def _get_document_info(doc: Document, prompt: BasePromptTemplate[str]) -> dict:405    base_info = {"page_content": doc.page_content, **doc.metadata}406    missing_metadata = set(prompt.input_variables).difference(base_info)407    if len(missing_metadata) > 0:408        required_metadata = [409            iv for iv in prompt.input_variables if iv != "page_content"410        ]411        msg = (412            f"Document prompt requires documents to have metadata variables: "413            f"{required_metadata}. Received document with missing metadata: "414            f"{list(missing_metadata)}."415        )416        raise ValueError(417            create_message(message=msg, error_code=ErrorCode.INVALID_PROMPT_INPUT)418        )419    return {k: base_info[k] for k in prompt.input_variables}420 421 422def format_document(doc: Document, prompt: BasePromptTemplate[str]) -> str:423    """Format a document into a string based on a prompt template.424 425    First, this pulls information from the document from two sources:426 427    1. `page_content`: This takes the information from the `document.page_content` and428        assigns it to a variable named `page_content`.429    2. `metadata`: This takes information from `document.metadata` and assigns it to430        variables of the same name.431 432    Those variables are then passed into the `prompt` to produce a formatted string.433 434    Args:435        doc: `Document`, the `page_content` and `metadata` will be used to create the436            final string.437        prompt: `BasePromptTemplate`, will be used to format the `page_content` and438            `metadata` into the final string.439 440    Returns:441        String of the document formatted.442 443    Example:444        ```python445        from langchain_core.documents import Document446        from langchain_core.prompts import PromptTemplate447 448        doc = Document(page_content="This is a joke", metadata={"page": "1"})449        prompt = PromptTemplate.from_template("Page {page}: {page_content}")450        format_document(doc, prompt)451        # -> "Page 1: This is a joke"452        ```453    """454    return prompt.format(**_get_document_info(doc, prompt))455 456 457async def aformat_document(doc: Document, prompt: BasePromptTemplate[str]) -> str:458    """Async format a document into a string based on a prompt template.459 460    First, this pulls information from the document from two sources:461 462    1. `page_content`: This takes the information from the `document.page_content` and463        assigns it to a variable named `page_content`.464    2. `metadata`: This takes information from `document.metadata` and assigns it to465        variables of the same name.466 467    Those variables are then passed into the `prompt` to produce a formatted string.468 469    Args:470        doc: `Document`, the `page_content` and `metadata` will be used to create the471            final string.472        prompt: `BasePromptTemplate`, will be used to format the `page_content` and473            `metadata` into the final string.474 475    Returns:476        String of the document formatted.477    """478    return await prompt.aformat(**_get_document_info(doc, prompt))479 
codekingpro/portable-devtools · Team Ai