codekingpro/portable-devtools
114k
1"""Dictionary prompt template."""2 3import warnings4from functools import cached_property5from typing import Any, Literal, cast6 7from pydantic import model_validator8from typing_extensions import override9 10from langchain_core.load import dumpd11from langchain_core.prompts.string import (12 DEFAULT_FORMATTER_MAPPING,13 get_template_variables,14)15from langchain_core.runnables import RunnableConfig, RunnableSerializable16from langchain_core.runnables.config import ensure_config17 18 19class DictPromptTemplate(RunnableSerializable[dict, dict]):20 """Template represented by a dictionary.21 22 Recognizes variables in f-string or mustache formatted string dict values.23 24 Does NOT recognize variables in dict keys. Applies recursively.25 26 Example:27 ```python28 prompt = DictPromptTemplate(29 template={30 "type": "text",31 "text": "Hello {name}",32 "metadata": {"source": "{source}"},33 },34 template_format="f-string",35 )36 prompt.format(name="Alice", source="docs")37 # {38 # "type": "text",39 # "text": "Hello Alice",40 # "metadata": {"source": "docs"},41 # }42 ```43 """44 45 template: dict[str, Any]46 template_format: Literal["f-string", "mustache"]47 48 @model_validator(mode="after")49 def validate_template(self) -> "DictPromptTemplate":50 """Validate that the template structure contains only safe variables."""51 _get_input_variables(self.template, self.template_format)52 return self53 54 @property55 def input_variables(self) -> list[str]:56 """Template input variables."""57 return _get_input_variables(self.template, self.template_format)58 59 def format(self, **kwargs: Any) -> dict[str, Any]:60 """Format the prompt with the inputs.61 62 Returns:63 A formatted dict.64 """65 return _insert_input_variables(self.template, kwargs, self.template_format)66 67 async def aformat(self, **kwargs: Any) -> dict[str, Any]:68 """Format the prompt with the inputs.69 70 Returns:71 A formatted dict.72 """73 return self.format(**kwargs)74 75 @override76 def invoke(77 self, input: dict, config: RunnableConfig | None = None, **kwargs: Any78 ) -> dict:79 return self._call_with_config(80 lambda x: self.format(**x),81 input,82 ensure_config(config),83 run_type="prompt",84 serialized=self._serialized,85 **kwargs,86 )87 88 @property89 def _prompt_type(self) -> str:90 return "dict-prompt"91 92 @cached_property93 def _serialized(self) -> dict[str, Any]:94 # self is always a Serializable object in this case, thus the result is95 # guaranteed to be a dict since dumpd uses the default callback, which uses96 # obj.to_json which always returns TypedDict subclasses97 return cast("dict[str, Any]", dumpd(self))98 99 @classmethod100 def is_lc_serializable(cls) -> bool:101 """Return `True` as this class is serializable."""102 return True103 104 @classmethod105 def get_lc_namespace(cls) -> list[str]:106 """Get the namespace of the LangChain object.107 108 Returns:109 `["langchain_core", "prompts", "dict"]`110 """111 return ["langchain_core", "prompts", "dict"]112 113 def pretty_repr(self, *, html: bool = False) -> str:114 """Human-readable representation.115 116 Args:117 html: Whether to format as HTML.118 119 Returns:120 Human-readable representation.121 """122 raise NotImplementedError123 124 125def _get_input_variables(126 template: dict, template_format: Literal["f-string", "mustache"]127) -> list[str]:128 input_variables = []129 for v in template.values():130 if isinstance(v, str):131 input_variables += get_template_variables(v, template_format)132 elif isinstance(v, dict):133 input_variables += _get_input_variables(v, template_format)134 elif isinstance(v, (list, tuple)):135 for x in v:136 if isinstance(x, str):137 input_variables += get_template_variables(x, template_format)138 elif isinstance(x, dict):139 input_variables += _get_input_variables(x, template_format)140 return list(set(input_variables))141 142 143def _insert_input_variables(144 template: dict[str, Any],145 inputs: dict[str, Any],146 template_format: Literal["f-string", "mustache"],147) -> dict[str, Any]:148 formatted: dict[str, Any] = {}149 formatter = DEFAULT_FORMATTER_MAPPING[template_format]150 for k, v in template.items():151 if isinstance(v, str):152 formatted[k] = formatter(v, **inputs)153 elif isinstance(v, dict):154 if k == "image_url" and "path" in v:155 msg = (156 "Specifying image inputs via file path in environments with "157 "user-input paths is a security vulnerability. Out of an abundance "158 "of caution, the utility has been removed to prevent possible "159 "misuse."160 )161 warnings.warn(msg, stacklevel=2)162 formatted[k] = _insert_input_variables(v, inputs, template_format)163 elif isinstance(v, (list, tuple)):164 formatted_v: list[str | dict[str, Any]] = []165 for x in v:166 if isinstance(x, str):167 formatted_v.append(formatter(x, **inputs))168 elif isinstance(x, dict):169 formatted_v.append(170 _insert_input_variables(x, inputs, template_format)171 )172 formatted[k] = type(v)(formatted_v)173 else:174 formatted[k] = v175 return formatted176 