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